{"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":"# Multimodal Single-Cell Integration Inference","metadata":{}},{"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\nfrom flax.serialization import from_bytes\n\nfrom typing import Sequence\n\nInputShapeCITE = 22085\nOutputCITE = 140\n\nInputShapeMULTI = 228942\nOutputMULTI = 23418\n\nmainUnitsMULTI  = [InputShapeMULTI//1024,InputShapeMULTI//2084,InputShapeMULTI//4168,OutputMULTI]\nmainUnitsCITE  = [InputShapeCITE,InputShapeCITE//512,InputShapeCITE//2048,OutputCITE]\nnp.random.seed(10)","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-11-15T21:30:09.920477Z","iopub.execute_input":"2022-11-15T21:30:09.921075Z","iopub.status.idle":"2022-11-15T21:30:12.337111Z","shell.execute_reply.started":"2022-11-15T21:30:09.920960Z","shell.execute_reply":"2022-11-15T21:30:12.335675Z"},"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 Regressor(nn.Module):\n    \n    Units: Sequence[int]\n    name: str \n    train: bool = True\n    \n    def setup(self):\n        self.main = Coder(self.Units[1::],self.name,self.train)\n        self.last = nn.Dense(self.Units[-1], name='output')\n    \n    @nn.compact\n    def __call__(self, inputs):\n        \n        x = inputs\n        mlp_main = self.main(x)\n        output = self.last(mlp_main)\n        \n        return output","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-11-15T21:30:12.339269Z","iopub.execute_input":"2022-11-15T21:30:12.339963Z","iopub.status.idle":"2022-11-15T21:30:12.355110Z","shell.execute_reply.started":"2022-11-15T21:30:12.339924Z","shell.execute_reply":"2022-11-15T21:30:12.353641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CITE Predictions","metadata":{}},{"cell_type":"code","source":"BasePathCITE = '../input/open-problems-single-cells/CITE/test-inputs'\ntestInputsCITE = [BasePathCITE+'/'+val for val in os.listdir(BasePathCITE)]\n\nModelsPathCITE = '../input/multimodal-singlecell-integration-models/CITE'\nModelsCITE = [ModelsPathCITE+'/'+val for val in os.listdir(ModelsPathCITE)]\n\ndef ModelCITE():\n    return Regressor(mainUnitsCITE,'test')","metadata":{"execution":{"iopub.status.busy":"2022-11-15T21:30:12.357158Z","iopub.execute_input":"2022-11-15T21:30:12.358175Z","iopub.status.idle":"2022-11-15T21:30:12.831016Z","shell.execute_reply.started":"2022-11-15T21:30:12.358082Z","shell.execute_reply":"2022-11-15T21:30:12.829586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for model in [ModelsCITE[1]]:\n    \n    with open(model, \"rb\") as state_f:\n        state = from_bytes(ModelCITE, state_f.read())\n        state = jax.tree_util.tree_map(jnp.array, state)\n    \n    localparams = {'params':state['params'],'batch_stats':state['batch_stats']}\n\n    def CITEModel(trainparams,batch):\n        return Regressor(mainUnitsCITE,'test',train=False).apply(trainparams, batch)\n    \n    ModelPredsCITE = []\n    Bsize = 10000\n    for k in range(0,len(testInputsCITE),Bsize):\n        localfilespath = testInputsCITE[k:k+Bsize]\n        localdata = [np.load(val) for val in localfilespath]\n        localdata = np.vstack(localdata)\n        preds = CITEModel(localparams,localdata)\n        ModelPredsCITE.append(np.array(preds))\n    \n    ModelPredsCITE = np.vstack(ModelPredsCITE)\n    del localdata","metadata":{"execution":{"iopub.status.busy":"2022-11-15T21:30:12.835001Z","iopub.execute_input":"2022-11-15T21:30:12.835580Z","iopub.status.idle":"2022-11-15T21:33:55.667376Z","shell.execute_reply.started":"2022-11-15T21:30:12.835526Z","shell.execute_reply":"2022-11-15T21:33:55.666024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MULTI Predictions","metadata":{}},{"cell_type":"code","source":"BasePathMULTI = '../input/open-problems-single-cells/MULTI/test-inputs'\ntestInputsMULTI = [BasePathMULTI+'/'+val for val in os.listdir(BasePathMULTI)]\n\nModelsPathMULTI = '../input/multimodal-singlecell-integration-models/MULTI'\nModelsMULTI = [ModelsPathMULTI+'/'+val for val in os.listdir(ModelsPathMULTI)]\n\ndef ModelMULTI():\n    return Regressor(mainUnitsMULTI,'test')","metadata":{"execution":{"iopub.status.busy":"2022-11-15T21:33:55.669256Z","iopub.execute_input":"2022-11-15T21:33:55.669755Z","iopub.status.idle":"2022-11-15T21:33:56.177160Z","shell.execute_reply.started":"2022-11-15T21:33:55.669705Z","shell.execute_reply":"2022-11-15T21:33:56.175692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for model in [ModelsMULTI[3]]:\n    \n    with open(model, \"rb\") as state_f:\n        state = from_bytes(ModelMULTI, state_f.read())\n        state = jax.tree_util.tree_map(jnp.array, state)\n    \n    localparams = {'params':state['params'],'batch_stats':state['batch_stats']}\n\n    def MULTIModel(trainparams,batch):\n        return Regressor(mainUnitsMULTI,'test',train=False).apply(trainparams, batch)\n    \n    ModelPredsMULTI = []\n    Bsize = 5000\n    for k in range(0,len(testInputsMULTI),Bsize):\n        localfilespath = testInputsMULTI[k:k+Bsize]\n        localdata = [np.load(val) for val in localfilespath]\n        localdata = np.vstack(localdata)\n        preds = MULTIModel(localparams,localdata)\n        ModelPredsMULTI.append(np.array(preds))\n    \n    ModelPredsMULTI = np.vstack(ModelPredsMULTI)\n    del localdata","metadata":{"execution":{"iopub.status.busy":"2022-11-15T21:33:56.178781Z","iopub.execute_input":"2022-11-15T21:33:56.179285Z","iopub.status.idle":"2022-11-15T21:52:55.404548Z","shell.execute_reply.started":"2022-11-15T21:33:56.179224Z","shell.execute_reply":"2022-11-15T21:52:55.402850Z"},"trusted":true},"execution_count":null,"outputs":[]}]}