{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "553c5749-a5b8-4f26-984f-0dfb3e69bb1e",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import zarr\n",
    "import yaml\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "from tqdm import tqdm\n",
    "\n",
    "def extract_data(loci_groups_df, data):\n",
    "    # Extract data for each loci group (500bp is better)\n",
    "    extracted_data = []\n",
    "    for index, row in tqdm(loci_groups_df.iterrows(), total=len(loci_groups_df)):\n",
    "        chr_name = row['seqname']\n",
    "        start = row['tss'] - 500\n",
    "        end = row['tss'] + 500\n",
    "        loci_data = data['chrs'][chr_name][start:end, :]\n",
    "        extracted_data.append((loci_data.max(axis=0), loci_data.mean(axis=0)))\n",
    "    if len(extracted_data) != 0:\n",
    "        extracted_data = np.array(extracted_data).transpose(0, 2, 1)\n",
    "    return extracted_data # Shape (n_loci, n_cap, 2[max, mean])\n",
    "\n",
    "def remove_chr_X(loci_groups_df):\n",
    "    loci_groups_df = loci_groups_df[loci_groups_df['seqname'] != 'chrX']\n",
    "    return loci_groups_df\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b8d55c16-3a0c-480c-b241-c3c0f829408c",
   "metadata": {},
   "outputs": [],
   "source": [
    "import re\n",
    "# Load config\n",
    "config_path='config_sene_vs_growing.yaml'\n",
    "with open('config_sene_vs_growing.yaml', 'r') as f:\n",
    "    config = yaml.safe_load(f)\n",
    "config_name=re.sub(r'^(?:config_)?|\\.ya?ml$', '',config_path)\n",
    "\n",
    "output_path = config['extract_data_config']['output']['path']+'/'+config_name\n",
    "os.makedirs(output_path, exist_ok=True)\n",
    "preprocessing_data_path = config['preprocessing_data_config']['output']['path']\n",
    "postprocess_data_path = config['merge_inference_config']['output']['path']\n",
    "\n",
    "experimental_celltype = config['loci_source_config']['atac']['data_path']['experimental']\n",
    "control_celltype = config['loci_source_config']['atac']['data_path']['control']\n",
    "\n",
    "# Load loci groups\n",
    "loci_groups_list = ['upregulated', 'downregulated', 'conserved']\n",
    "loci_groups_df_dict = {}\n",
    "for loci_group in loci_groups_list:\n",
    "    loci_groups_df = pd.read_csv(os.path.join(preprocessing_data_path,config_name, 'gene_list', f'{loci_group}_gene_list.csv'))\n",
    "    loci_groups_df = remove_chr_X(loci_groups_df)\n",
    "    loci_groups_df_dict[loci_group] = loci_groups_df\n",
    "    # Save loci groups\n",
    "    loci_groups_df.to_csv(os.path.join(output_path,f'{loci_group}_gene_list_no_chrX.csv'), index=False)\n",
    "\n",
    "# Load inference data\n",
    "celltypes = experimental_celltype.split(',')+ control_celltype.split(',')\n",
    "data_dict = {}\n",
    "for celltype in celltypes:\n",
    "    data_path = os.path.join(postprocess_data_path, celltype, 'data.zarr')\n",
    "    data = zarr.open(data_path, mode='r')\n",
    "    data_dict[celltype] = data\n",
    "\n",
    "# Extract data for each loci group (500bp is better) for each celltype\n",
    "for loci_group in loci_groups_list:\n",
    "    data_celltype_list = []\n",
    "    for celltype in celltypes:\n",
    "        loci_groups_df = loci_groups_df_dict[loci_group]\n",
    "        cell_type_loci_data = extract_data(loci_groups_df, data_dict[celltype])\n",
    "        data_celltype_list.append(cell_type_loci_data) # Shape (n_loci, n_cap, 2[max, mean])\n",
    "    data_loci = np.array(data_celltype_list) # Shape (n_celltypes, n_loci, n_cap, 2[max, mean])\n",
    "\n",
    "    # Save data and annotations\n",
    "    np.save(os.path.join(output_path, f'{loci_group}_data.npy'), data_loci)\n",
    "\n",
    "    # Save cell types, loci groups, and loci\n",
    "    with open(os.path.join(output_path, f'{loci_group}_celltypes.txt'), 'w') as f:\n",
    "        for celltype in celltypes:\n",
    "            f.write(celltype + '\\n')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b42671cc-b1f9-42a8-aff3-a083ff9055a3",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "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
}
