{"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":"# Dimensionality reduction with VAEs\n\nLarge datasets always offer some challenges. And those can increase if the resources appear to not be enough to analyze them. A quick and simple solution to work with large datasets is to split them into small files. Then process them by batch. The following describes a simple VAE for dimensionality reduction using the dataset from  Open Problems - Multimodal Single-Cell Integration Predict how DNA, RNA & protein measurements co-vary in single cells. ","metadata":{}},{"cell_type":"markdown","source":"## Dataset generation\n\nFiles are loaded in batches and saved into single-cell files. The name of the file is equal to the cell id. Thus the cell-id can be used to define the different folds for cross-validation. A similar approach can be used to train the models by loading a batch of files. ","metadata":{}},{"cell_type":"code","source":"'''\nnromws = 120000\nstep = 10000\nfor k in range(0,nromws,step):\n    localdf = pd.read_hdf(inputFile,start=k,stop=k+step)\n    localIndex = localdf.index\n            \n    for inx in localIndex:\n        rowval = localdf.loc[inx]\n        row = np.array(rowval).astype(typs)\n        np.save(OutDir+'/'+inx+'.npy',row)\n'''","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport 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 KFold\nfrom sklearn.model_selection import train_test_split\n\nimport jax\nimport optax\nimport jax.numpy as jnp\nfrom jax import random\nfrom flax import linen as nn\n\nfrom typing import Sequence\n\nInputShapeCITE = 22085\nInputShapeMULTI = 228942\n\nmainUnitsCITE  = [InputShapeCITE,InputShapeCITE//512,InputShapeCITE//2048,2]\nmainUnitsMULTI  = [InputShapeMULTI,InputShapeMULTI//1024,InputShapeMULTI//8192,2]\n\nsh = 0.0001\nbatchSize = 128","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-30T09:23:06.557548Z","iopub.execute_input":"2022-10-30T09:23:06.557896Z","iopub.status.idle":"2022-10-30T09:23:06.567417Z","shell.execute_reply.started":"2022-10-30T09:23:06.557870Z","shell.execute_reply":"2022-10-30T09:23:06.566025Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-30T08:38:20.767771Z","iopub.execute_input":"2022-10-30T08:38:20.768130Z","iopub.status.idle":"2022-10-30T08:38:20.790757Z","shell.execute_reply.started":"2022-10-30T08:38:20.768101Z","shell.execute_reply":"2022-10-30T08:38:20.789478Z"},"jupyter":{"source_hidden":true},"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,basePath,dataShape,epochs=10,batch_size=64,lr=0.005):\n    \n    localOptimizer = optax.adam(learning_rate=lr)\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            batchNames = TrainData[k:k+batch_size]\n            \n            batch = np.zeros(shape=(len(batchNames),dataShape)) \n            \n            for ii,val in enumerate(batchNames):\n                 batch[ii,:] = np.load(basePath+'/'+val+'.npy')\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_batchNames = TestData[i:i+batch_size]\n            \n            val_batch = np.zeros(shape=(len(val_batchNames),dataShape)) \n            \n            for ii,val in enumerate(val_batchNames):\n                 val_batch[ii,:] = np.load(basePath+'/'+val+'.npy')\n            \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","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-30T08:38:20.792685Z","iopub.execute_input":"2022-10-30T08:38:20.793044Z","iopub.status.idle":"2022-10-30T08:38:20.812648Z","shell.execute_reply.started":"2022-10-30T08:38:20.793019Z","shell.execute_reply":"2022-10-30T08:38:20.811385Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CITE Dataset\n## Loading the data","metadata":{}},{"cell_type":"code","source":"BasePath = '../input/open-problems-single-cells/CITE'\ntrainInputsCITE = BasePath + '/train-inputs'\n\nfiles_inputs = os.listdir(trainInputsCITE)\nvalid_cells = [val[0:-4] for val in files_inputs]\nnp.random.shuffle(valid_cells)\n\ntrainCells, valCells, _, _ = train_test_split(valid_cells,np.arange(len(valid_cells)) , test_size=0.1, random_state=42)\n\nXtrain, Xtest, _, _ = train_test_split(trainCells,np.arange(len(trainCells)) , test_size=0.1, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2022-10-30T08:38:20.835205Z","iopub.execute_input":"2022-10-30T08:38:20.835502Z","iopub.status.idle":"2022-10-30T08:38:20.922434Z","shell.execute_reply.started":"2022-10-30T08:38:20.835479Z","shell.execute_reply":"2022-10-30T08:38:20.921246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Training ","metadata":{}},{"cell_type":"code","source":"def VAEModel():\n    return VAE(mainUnitsCITE,'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, InputShapeCITE), jnp.float32)\ninitModel = VAEModel().init(key, init_data, rng)\n    \nparams0 = initModel['params']\nbatchStats = initModel['batch_stats']\n\ntrloss,tstloss,params0,batchStats = TrainModel(Xtrain,Xtest,loss,params0,batchStats,rng,trainInputsCITE,InputShapeCITE,lr=0.015,epochs=15,batch_size=128)\n    \nfinalParams = {'params':params0,'batch_stats':batchStats}","metadata":{"execution":{"iopub.status.busy":"2022-10-30T08:38:20.925522Z","iopub.execute_input":"2022-10-30T08:38:20.925935Z","iopub.status.idle":"2022-10-30T09:02:34.124072Z","shell.execute_reply.started":"2022-10-30T08:38:20.925908Z","shell.execute_reply":"2022-10-30T09:02:34.123172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Latent space visualization","metadata":{}},{"cell_type":"code","source":"def EncoderModel(trainparams,batch):\n    return Encoder(mainUnitsCITE,'testencoder',train=False).apply(trainparams, batch)\n    \nlocalparams = {'params':finalParams['params']['testencoder'],'batch_stats':finalParams['batch_stats']['testencoder']}\n\nlocalData = []\n\nfor val in valCells:\n    localData.append(np.load(trainInputsCITE+'/'+val+'.npy'))\n\nlocalData = np.array(localData)\n   \nmu,logvar = EncoderModel(localparams,localData)\nVariationalRepresentation = reparameterize(rng,mu,logvar)","metadata":{"execution":{"iopub.status.busy":"2022-10-30T09:02:34.125555Z","iopub.execute_input":"2022-10-30T09:02:34.126085Z","iopub.status.idle":"2022-10-30T09:02:42.734411Z","shell.execute_reply.started":"2022-10-30T09:02:34.126051Z","shell.execute_reply":"2022-10-30T09:02:42.732498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.scatter(VariationalRepresentation[:,0],VariationalRepresentation[:,1],alpha=0.1)","metadata":{"execution":{"iopub.status.busy":"2022-10-30T09:02:42.736100Z","iopub.execute_input":"2022-10-30T09:02:42.736462Z","iopub.status.idle":"2022-10-30T09:02:42.960632Z","shell.execute_reply.started":"2022-10-30T09:02:42.736427Z","shell.execute_reply":"2022-10-30T09:02:42.959613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MULTI Dataset\n## Loading the data","metadata":{}},{"cell_type":"code","source":"BasePathMULTI = '../input/open-problems-single-cells/MULTI'\ntrainInputsMULTI = BasePathMULTI + '/train-inputs'\n\nfiles_inputsMULTI = os.listdir(trainInputsMULTI)\nvalid_cellsMULTI = [val[0:-4] for val in files_inputsMULTI]\nnp.random.shuffle(valid_cellsMULTI)\n\ntrainCellsMULTI, valCellsMULTI, _, _ = train_test_split(valid_cellsMULTI,np.arange(len(valid_cellsMULTI)) , test_size=0.1, random_state=42)\nXtrainMULTI, XtestMULTI, _, _ = train_test_split(trainCellsMULTI,np.arange(len(trainCellsMULTI)) , test_size=0.1, random_state=42)\n","metadata":{"execution":{"iopub.status.busy":"2022-10-30T09:23:17.300280Z","iopub.execute_input":"2022-10-30T09:23:17.300644Z","iopub.status.idle":"2022-10-30T09:23:17.409755Z","shell.execute_reply.started":"2022-10-30T09:23:17.300619Z","shell.execute_reply":"2022-10-30T09:23:17.408503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training the model","metadata":{}},{"cell_type":"code","source":"def VAEModel():\n    return VAE(mainUnitsMULTI,'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, InputShapeMULTI), jnp.float32)\ninitModel = VAEModel().init(key, init_data, rng)\n    \nparams0 = initModel['params']\nbatchStats = initModel['batch_stats']\n\ntrloss,tstloss,params0,batchStats = TrainModel(XtrainMULTI,XtestMULTI,loss,params0,batchStats,rng,trainInputsMULTI,InputShapeMULTI,lr=0.015,epochs=5,batch_size=128)\n    \nfinalParams = {'params':params0,'batch_stats':batchStats}","metadata":{"execution":{"iopub.status.busy":"2022-10-30T09:23:19.903205Z","iopub.execute_input":"2022-10-30T09:23:19.903548Z","iopub.status.idle":"2022-10-30T09:42:46.761757Z","shell.execute_reply.started":"2022-10-30T09:23:19.903521Z","shell.execute_reply":"2022-10-30T09:42:46.759792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Latent space visualization","metadata":{}},{"cell_type":"code","source":"def EncoderModel(trainparams,batch):\n    return Encoder(mainUnitsMULTI,'testencoder',train=False).apply(trainparams, batch)\n    \nlocalparams = {'params':finalParams['params']['testencoder'],'batch_stats':finalParams['batch_stats']['testencoder']}\n\nlocalData = []\n\nfor val in valCellsMULTI:\n    localData.append(np.load(trainInputsMULTI+'/'+val+'.npy'))\n\nlocalData = np.array(localData)\n   \nmu,logvar = EncoderModel(localparams,localData)\nVariationalRepresentation = reparameterize(rng,mu,logvar)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.scatter(VariationalRepresentation[:,0],VariationalRepresentation[:,1],alpha=0.1)","metadata":{},"execution_count":null,"outputs":[]}]}