{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "405fd37f-e105-4c3d-966e-9c8ce409b104",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "import yaml\n",
    "import matplotlib.pyplot as plt\n",
    "import matplotlib as mpl\n",
    "import json\n",
    "import argparse\n",
    "import re\n",
    "from src.signature import kmeans, nmf\n",
    "from src.utils.vis import figure, subplots\n",
    "from src.utils.signature import plot_signature_heatmap_full, plot_signature_heatmap_group\n",
    "\n",
    "\n",
    "    # Load config\n",
    "\n",
    "config_path='config_sene_vs_growing.yaml'\n",
    "# resolution order: CLI --config (in order) > env APP_CONFIG > ./config.yaml\n",
    "\n",
    "config_name=re.sub(r'^(?:config_)?|\\.ya?ml$', '',config_path)\n",
    "with open(config_path, 'r') as f:\n",
    "    config = yaml.safe_load(f)\n",
    "\n",
    "signature_method = config['signature_config']['method']\n",
    "print(f'Using {signature_method} method for signature extraction')\n",
    "postprocess_folder = config['postprocess_config']['output']['path']+'/'+config_name #hua\n",
    "\n",
    "output_path = os.path.join(config['signature_config']['output']['path'], signature_method,config_name) #hua\n",
    "os.makedirs(output_path, exist_ok=True)\n",
    "\n",
    "data_types = ['upregulated', 'downregulated']\n",
    "\n",
    "data_type='upregulated'\n",
    "print(f\"Processing {data_type} data\")\n",
    "\n",
    "data_type_folder = os.path.join(postprocess_folder, data_type)\n",
    "metadata_path = os.path.join(data_type_folder, 'metadata.json')\n",
    "with open(metadata_path, 'r') as f:\n",
    "    metadata = json.load(f)\n",
    "\n",
    "data_save_path = os.path.join(output_path, data_type)\n",
    "os.makedirs(data_save_path, exist_ok=True)\n",
    "\n",
    "# Find loci/CAP group signatures using method of choice\n",
    "plot_dict = {}\n",
    "\n",
    "\n",
    "\n",
    "celltypes = [metadata['control_name'], metadata['experiment_name']]\n",
    "for celltype in celltypes:\n",
    "print(f'Processing celltype: {celltype}')\n",
    "celltype_mat_path = os.path.join(data_type_folder, f'{celltype}.npy')\n",
    "celltype_mat = np.load(celltype_mat_path)\n",
    "\n",
    "celltype_save_path = os.path.join(data_save_path, celltype)\n",
    "os.makedirs(celltype_save_path, exist_ok=True)\n",
    "\n",
    "plot_dict[celltype] = nmf.get_signature(\n",
    "                    data_save_path=celltype_save_path,  \n",
    "                    data_mat=celltype_mat,\n",
    "                    loci_list=metadata['filtered_gene_names'],\n",
    "                    cap_list=metadata['filtered_cap_names'],\n",
    "                    direction=data_type,\n",
    "                    rerun_cluster=True\n",
    "                )\n",
    "plot_signature_heatmap_full(plot_dict[celltype], block_size=400, vmin=-10, vmax=10, show_ticks=True)\n",
    "plot_signature_heatmap_group(plot_dict[celltype])\n",
    "\n",
    "# Compute similarity between celltype groups\n",
    "nmf.compute_group_similarity(plot_dict[celltypes[0]], plot_dict[celltypes[1]], celltypes, data_save_path)\n",
    "\n"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.11.13"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
