{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e2dc55e3-7fc8-4dcd-8121-a50170d292b5",
   "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",
    "from src.signature.nmf import *\n",
    "# Load config\n",
    "#parser = argparse.ArgumentParser() #hua added for differet configs\n",
    "\n",
    "#args = parser.parse_args()\n",
    "\n",
    "# resolution order: CLI --config (in order) > env APP_CONFIG > ./config.yaml\n",
    "#config_path = args.config\n",
    "config_path='config_sene_vs_growing.yaml'\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",
    "\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",
    "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",
    "celltypes = [metadata['control_name'], metadata['experiment_name']]\n",
    "#for celltype in celltypes:\n",
    "celltype='sene'\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",
    "'''\n",
    "\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",
    "cluster_loci = True \n",
    "cluster_cap = True\n",
    "rerun_cluster = False\n",
    "\n",
    "# Load config\n",
    "with open('config.yaml', 'r') as f:\n",
    "    config = yaml.safe_load(f)\n",
    "\n",
    "# NMF parameters from config\n",
    "random_state = config['signature_config']['seed']\n",
    "n_components = config['signature_config']['nmf']['n_components']\n",
    "init = config['signature_config']['nmf']['init']\n",
    "max_iter = config['signature_config']['nmf']['max_iter']\n",
    "n_components_range = range(config['signature_config']['nmf']['min_components'], config['signature_config']['nmf']['max_components'] + 1)\n",
    "\n",
    "W_npy_path = os.path.join(data_save_path, 'W_loci_by_k.npy')\n",
    "H_npy_path = os.path.join(data_save_path, 'H_k_by_cap.npy')\n",
    "W_df_path = os.path.join(data_save_path, 'W_loci_by_k.csv')\n",
    "H_df_path = os.path.join(data_save_path, 'H_k_by_cap.csv')\n",
    "\n",
    "def run_nmf(data, n_components, init='random', random_state=9, max_iter=200):\n",
    "    print(\"Running vanilla NMF with n_components =\", n_components)\n",
    "    from sklearn.decomposition import NMF\n",
    "    model = NMF(n_components=n_components, init=init, random_state=random_state, max_iter=max_iter)\n",
    "    W = model.fit_transform(data)\n",
    "    H = model.components_\n",
    "    return W, H\n",
    "\n",
    "\n",
    "W, H = run_nmf(data=data_mat, n_components=n_components, init=init, random_state=random_state, max_iter=max_iter)\n",
    "\n",
    "# Save W and H matrices\n",
    "np.save(W_npy_path, W)\n",
    "np.save(H_npy_path, H)\n",
    "W_df = pd.DataFrame(W, index=loci_list)\n",
    "W_df.to_csv(W_df_path, index=True)\n",
    "H_df = pd.DataFrame(H, columns=cap_list)\n",
    "H_df.to_csv(H_df_path, index=True)\n",
    "\n",
    "# Group CAPs and loci\n",
    "loci_group = group_loci(W_df, data_mat)\n",
    "cap_group = group_cap(H_df, data_mat)\n",
    "loci_group['group'] = loci_group['group'].astype(int)\n",
    "cap_group['group'] = cap_group['group'].astype(int)\n",
    "\n",
    "# Refine ranks of loci and CAPs within each group\n",
    "loci_rank, cap_rank = iterative_refinement(data_mat, loci_group['group'].values, cap_group['group'].values, loci_list, cap_list, top_group_percent = 0.5, top_item_percent = 0.5)\n",
    "loci_group['rank'] = loci_rank\n",
    "cap_group['rank'] = cap_rank\n",
    "\n",
    "# Build loci_group and cap_group dataframes\n",
    "loci_assignment = loci_group['group'].values.astype(int)\n",
    "cap_assignment = cap_group['group'].values.astype(int)\n",
    "loci_group = pd.DataFrame({'loci': loci_list, 'group': loci_assignment, 'avg_binding': np.mean(data_mat, axis=1), 'rank': loci_rank})\n",
    "cap_group = pd.DataFrame({'cap': cap_list, 'group': cap_assignment, 'avg_binding': np.mean(data_mat, axis=0), 'rank': cap_rank})\n",
    "\n",
    "# Track original positions\n",
    "loci_group['original_index'] = range(len(loci_group))\n",
    "cap_group['original_index'] = range(len(cap_group))\n",
    "\n",
    "# Re-arrange the loci and caps\n",
    "loci_group = loci_group.sort_values(['group', 'rank'], ascending=[True, True])\n",
    "cap_group = cap_group.sort_values(['group', 'rank'], ascending=[True, True])\n",
    "\n",
    "# Build the plot matrix\n",
    "plot_mat = data_mat[np.ix_(loci_group['original_index'], cap_group['original_index'])]\n",
    "loci_group = loci_group.reset_index(drop=True)\n",
    "cap_group = cap_group.reset_index(drop=True)\n",
    "\n",
    "# Save loci_group and cap_group\n",
    "loci_group.drop(columns=['original_index']).to_csv(os.path.join(data_save_path, 'loci_group.csv'))\n",
    "cap_group.drop(columns=['original_index']).to_csv(os.path.join(data_save_path, 'cap_group.csv'))\n",
    "# Save plot_mat\n",
    "plot_mat_path = os.path.join(data_save_path, 'plot_matrix.npy')\n",
    "np.save(plot_mat_path, plot_mat)\n",
    "\n",
    "# Plot group sizes\n",
    "plot_group_sizes(loci_group, cap_group, data_save_path)\n",
    "\n",
    "# Save group lists in FASTA-like format\n",
    "save_group_lists_fasta(loci_group, cap_group, data_save_path)\n",
    "\n",
    "# For plotting heatmap\n",
    "x_name = 'CAP'\n",
    "y_name = 'Loci'\n",
    "x_annotation = cap_group['cap']\n",
    "y_annotation = loci_group['loci'].values\n",
    "title_str = 'Average max CAP binding at loci'\n",
    "\n",
    "loci_k = n_components\n",
    "cap_k = n_components\n",
    "\n",
    "plot_dict = {\n",
    "    'plot_mat': plot_mat,\n",
    "    'loci_k': loci_k,\n",
    "    'cap_k': cap_k,\n",
    "    'loci_group': loci_group,\n",
    "    'cap_group': cap_group,\n",
    "    'data_save_path': data_save_path,\n",
    "    'x_name': x_name,\n",
    "    'y_name': y_name,\n",
    "    'x_annotation': x_annotation,\n",
    "    'y_annotation': y_annotation,\n",
    "    'title_str': title_str,\n",
    "    'direction': direction\n",
    "}\n",
    "\n",
    "return plot_dict\n",
    "\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",
    "\n",
    "    \n",
    "plot_signature_heatmap_full(plot_dict[celltype], block_size=400, vmin=-10, vmax=10, show_ticks=True)\n",
    "plot_dict=plot_dict[celltype]\n",
    "block_size=400\n",
    "vmin=-10\n",
    "vmax=10\n",
    "show_ticks=True\n",
    "\n",
    "plot_mat = plot_dict['plot_mat']\n",
    "loci_k = plot_dict['loci_k']\n",
    "cap_k = plot_dict['cap_k']\n",
    "loci_group = plot_dict['loci_group']\n",
    "cap_group = plot_dict['cap_group']\n",
    "data_save_path = plot_dict['data_save_path']\n",
    "x_name = plot_dict['x_name']\n",
    "y_name = plot_dict['y_name']\n",
    "x_annotation = plot_dict['x_annotation']\n",
    "y_annotation = plot_dict['y_annotation']\n",
    "title_str = plot_dict['title_str']\n",
    "direction = plot_dict['direction']\n",
    "def invert_heatmap(plot_mat, loci_k, cap_k, loci_group, cap_group, x_name, y_name, x_annotation, y_annotation):\n",
    "    return plot_mat.T, cap_k, loci_k, cap_group, loci_group, y_name, x_name, y_annotation, x_annotation\n",
    "\n",
    "plot_mat, loci_k, cap_k, loci_group, cap_group, x_name, y_name, x_annotation, y_annotation = invert_heatmap(plot_mat, loci_k, cap_k, loci_group, cap_group, x_name, y_name, x_annotation, y_annotation)\n",
    "\n",
    "save_name = f'{data_save_path}/signature_heatmap_full.pdf'\n",
    "#width = len(cap_group) * 1\n",
    "#height = len(loci_group) * 1\n",
    "width, height = 45, 45\n",
    "fig, ax = subplots(figsize=(width, height))\n",
    "\n",
    "#ax.imshow(plot_mat, cmap='coolwarm', aspect='auto', vmin=-10, vmax=10, interpolation='None', rasterized=False)\n",
    "# Instead use pcolormesh to plot heatmap\n",
    "if direction == 'upregulated':\n",
    "    plt.gca().invert_yaxis() # This is to invert the y-axis for down heatmap\n",
    "#ax.pcolormesh(plot_mat, cmap='coolwarm', vmin=-10, vmax=10)\n",
    "#ax.pcolormesh(plot_mat, cmap='coolwarm', vmin=-10, vmax=10)\n",
    "# Draw lines to separate groups\n",
    "from skimage.measure import block_reduce\n",
    "width, height = plot_mat.shape\n",
    "x_block_size = max(width // block_size, 1)\n",
    "y_block_size = max(height // block_size, 1)\n",
    "heatmap = block_reduce(plot_mat, (x_block_size, y_block_size), np.max)\n",
    "\n",
    "yedges = np.linspace(0, 1, heatmap.shape[0])\n",
    "xedges = np.linspace(0, 1, heatmap.shape[1])\n",
    "#ax.imshow(heatmap, cmap='viridis', aspect='auto', interpolation='none', rasterized=False, vmin=0, vmax=2)\n",
    "if vmin is None:\n",
    "    vmin = np.min(heatmap)\n",
    "if vmax is None:\n",
    "    vmax = np.max(heatmap)\n",
    "ax.pcolormesh(xedges, yedges, heatmap, cmap='RdBu_r', rasterized=False, vmin=vmin, vmax=vmax)\n",
    "#ax.imshow(heatmap, cmap='RdBu_r', aspect='auto', interpolation='none', rasterized=False, vmin=-10, vmax=10)\n",
    "width_block = width // block_size\n",
    "height_block = height // block_size\n",
    "\n",
    "y_delta = yedges[1] - yedges[0]\n",
    "x_delta = xedges[1] - xedges[0]\n",
    "\n",
    "for i in range(0, loci_k):\n",
    "    group_indices = np.where(loci_group['group'] == i)[0]\n",
    "    if len(group_indices) > 0:  # Check if group exists and is not empty\n",
    "        h_idx = group_indices[-1] / width\n",
    "        # Select the closest yedge\n",
    "        h_idx = yedges[np.argmin(np.abs(yedges - h_idx))]\n",
    "        h_idx = h_idx + y_delta / 2\n",
    "        ax.axhline(h_idx, color='black', linewidth=0.2, linestyle='--', alpha = 0.5)\n",
    "for i in range(0, cap_k):\n",
    "    group_indices = np.where(cap_group['group'] == i)[0]\n",
    "    if len(group_indices) > 0:  # Check if group exists and is not empty\n",
    "        v_idx = group_indices[-1] / height\n",
    "        # Select the closest xedge\n",
    "        v_idx = np.argmin(np.abs(xedges - v_idx))\n",
    "        v_idx = xedges[v_idx] - x_delta / 2\n",
    "        ax.axvline(v_idx, color='black', linewidth=0.2, linestyle='--', alpha = 0.5)\n",
    "'''\n",
    "ax.set_xlabel('Position')\n",
    "ax.set_ylabel('CAP')\n",
    "'''\n",
    "if show_ticks:\n",
    "    total_ticks = 10\n",
    "    '''\n",
    "    ax.set_yticks(yedges[::len(yedges) // total_ticks][:total_ticks])\n",
    "    ax.set_xticks(xedges[::len(xedges) // total_ticks][:total_ticks])\n",
    "    ax.set_yticklabels(range(width)[::width // total_ticks][:total_ticks])\n",
    "    ax.set_xticklabels(range(height)[::height // total_ticks][:total_ticks])\n",
    "    # Rotate x tick labels\n",
    "    ax.set_xticklabels(ax.get_xticklabels(), rotation=90)\n",
    "    '''\n",
    "    ax.xaxis.set_major_locator(mticker.MaxNLocator(nbins=total_ticks))\n",
    "    ax.yaxis.set_major_locator(mticker.MaxNLocator(nbins=total_ticks)) \n",
    "    ax.set_xticklabels(ax.get_xticklabels(), rotation=90)\n",
    "else:\n",
    "    ax.set_xticks([])\n",
    "    ax.set_yticks([])\n",
    "# remove x spine\n",
    "ax.spines['top'].set_visible(True)\n",
    "ax.spines['right'].set_visible(True)\n",
    "ax.set_xlabel(x_name)\n",
    "ax.set_ylabel(y_name)\n",
    "#ax.set_xticks(np.arange(len(cap_group)) + 0.5)\n",
    "#ax.set_yticks(np.arange(len(loci_group)) + 0.5)\n",
    "#ax.set_xticklabels(x_annotation, rotation=90)\n",
    "#ax.set_yticklabels(y_annotation)\n",
    "# Change label font size\n",
    "#ax.xaxis.set_tick_params(labelsize=2)\n",
    "#ax.yaxis.set_tick_params(labelsize=2)\n",
    "ax.set_title(title_str)\n",
    "plt.savefig(save_name, bbox_inches='tight')\n",
    "plt.close()\n",
    "\n",
    "def plot_signature_heatmap_group(plot_dict):\n",
    "    plot_mat = plot_dict['plot_mat']\n",
    "    loci_k = plot_dict['loci_k']\n",
    "    cap_k = plot_dict['cap_k']\n",
    "    loci_group = plot_dict['loci_group']\n",
    "    cap_group = plot_dict['cap_group']\n",
    "    data_save_path = plot_dict['data_save_path']\n",
    "    x_name = plot_dict['x_name']\n",
    "    y_name = plot_dict['y_name']\n",
    "    x_annotation = plot_dict['x_annotation']    \n",
    "    y_annotation = plot_dict['y_annotation']\n",
    "    title_str = plot_dict['title_str']\n",
    "    \n",
    "    # Plot heatmap of groups\n",
    "    group_mean_mat = np.zeros((loci_k, cap_k))\n",
    "    for i in range(loci_k):\n",
    "        for j in range(cap_k):\n",
    "            loci_mask = loci_group['group'] == i\n",
    "            cap_mask = cap_group['group'] == j\n",
    "            \n",
    "            # Check if both groups have members\n",
    "            if np.any(loci_mask) and np.any(cap_mask):\n",
    "                group_mean_mat[i, j] = np.mean(plot_mat[loci_mask, :][:, cap_mask])\n",
    "            else:\n",
    "                group_mean_mat[i, j] = 0  # Set to 0 if either group is empty\n",
    "    group_mean_df = pd.DataFrame(\n",
    "        group_mean_mat,\n",
    "        index = [f'loci_group_{i}' for i in range(loci_k)], \n",
    "        columns = [f'cap_group_{i}' for i in range(cap_k)]\n",
    "        )\n",
    "    group_mean_df.to_csv(f'{data_save_path}/group_mean_matrix.csv')\n",
    "    save_name = f'{data_save_path}/signature_heatmap_group.pdf'\n",
    "\n",
    "    plot_mat, loci_k, cap_k, loci_group, cap_group, x_name, y_name, x_annotation, y_annotation = invert_heatmap(group_mean_mat, loci_k, cap_k, loci_group, cap_group, x_name, y_name, x_annotation, y_annotation)\n",
    "    fig, ax = subplots(figsize=(70, 70))\n",
    "    ax.imshow(plot_mat, cmap='coolwarm', aspect='auto', vmin=-5, vmax=5, interpolation=None)\n",
    "    ax.set_xlabel(x_name)\n",
    "    ax.set_ylabel(y_name)\n",
    "    ax.set_xticks(np.arange(cap_k))\n",
    "    ax.set_yticks(np.arange(loci_k))\n",
    "    ax.set_xticklabels(np.arange(cap_k))\n",
    "    ax.set_yticklabels(np.arange(loci_k))\n",
    "    ax.set_title(title_str)\n",
    "    plt.savefig(save_name)\n",
    "    plt.close()\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
}
