{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "e5df1bf1-6066-46ac-9f58-5d9eba817c65",
   "metadata": {},
   "outputs": [
    {
     "ename": "KeyboardInterrupt",
     "evalue": "",
     "output_type": "error",
     "traceback": [
      "\u001b[31m---------------------------------------------------------------------------\u001b[39m",
      "\u001b[31mKeyboardInterrupt\u001b[39m                         Traceback (most recent call last)",
      "\u001b[36mCell\u001b[39m\u001b[36m \u001b[39m\u001b[32mIn[1]\u001b[39m\u001b[32m, line 5\u001b[39m\n\u001b[32m      2\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mos\u001b[39;00m \n\u001b[32m      3\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mnumpy\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mnp\u001b[39;00m\n\u001b[32m----> \u001b[39m\u001b[32m5\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01msrc\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mChromnitronDataset\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m ChromnitronDataset, create_from_original_data\n\u001b[32m      6\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01msrc\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mutils\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mvis\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m plot_histogram, plot_chipseq_tracks, plot_variability \n\u001b[32m      7\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mre\u001b[39;00m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m/gpfs/scratch/zhouh05/aris/senescence_202411/downstream/src/ChromnitronDataset.py:5\u001b[39m\n\u001b[32m      3\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mnumpy\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mnp\u001b[39;00m\n\u001b[32m      4\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mpandas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mpd\u001b[39;00m\n\u001b[32m----> \u001b[39m\u001b[32m5\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mpyplot\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mplt\u001b[39;00m\n\u001b[32m      6\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mtyping\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m List, Optional, Tuple\n\u001b[32m      8\u001b[39m \u001b[38;5;28;01mclass\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mChromnitronDataset\u001b[39;00m:\n",
      "\u001b[36mFile \u001b[39m\u001b[32m~/.conda/envs/chromni/lib/python3.11/site-packages/matplotlib/__init__.py:161\u001b[39m\n\u001b[32m    157\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mpackaging\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mversion\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m parse \u001b[38;5;28;01mas\u001b[39;00m parse_version\n\u001b[32m    159\u001b[39m \u001b[38;5;66;03m# cbook must import matplotlib only within function\u001b[39;00m\n\u001b[32m    160\u001b[39m \u001b[38;5;66;03m# definitions, so it is safe to import from it here.\u001b[39;00m\n\u001b[32m--> \u001b[39m\u001b[32m161\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01m.\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m _api, _version, cbook, _docstring, rcsetup\n\u001b[32m    162\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01m_api\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m MatplotlibDeprecationWarning\n\u001b[32m    163\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mrcsetup\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m cycler  \u001b[38;5;66;03m# noqa: F401\u001b[39;00m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m~/.conda/envs/chromni/lib/python3.11/site-packages/matplotlib/rcsetup.py:28\u001b[39m\n\u001b[32m     26\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mbackends\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m BackendFilter, backend_registry\n\u001b[32m     27\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mcbook\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m ls_mapper\n\u001b[32m---> \u001b[39m\u001b[32m28\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mcolors\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m Colormap, is_color_like\n\u001b[32m     29\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01m_fontconfig_pattern\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m parse_fontconfig_pattern\n\u001b[32m     30\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01m_enums\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m JoinStyle, CapStyle\n",
      "\u001b[36mFile \u001b[39m\u001b[32m~/.conda/envs/chromni/lib/python3.11/site-packages/matplotlib/colors.py:53\u001b[39m\n\u001b[32m     50\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mre\u001b[39;00m\n\u001b[32m     52\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mPIL\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m Image\n\u001b[32m---> \u001b[39m\u001b[32m53\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mPIL\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mPngImagePlugin\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m PngInfo\n\u001b[32m     55\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmatplotlib\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mmpl\u001b[39;00m\n\u001b[32m     56\u001b[39m \u001b[38;5;28;01mimport\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mnumpy\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mas\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mnp\u001b[39;00m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m~/.conda/envs/chromni/lib/python3.11/site-packages/PIL/PngImagePlugin.py:45\u001b[39m\n\u001b[32m     42\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01menum\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m IntEnum\n\u001b[32m     43\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mtyping\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m IO, Any, NamedTuple, NoReturn, cast\n\u001b[32m---> \u001b[39m\u001b[32m45\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01m.\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m Image, ImageChops, ImageFile, ImagePalette, ImageSequence\n\u001b[32m     46\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01m.\u001b[39;00m\u001b[34;01m_binary\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m i16be \u001b[38;5;28;01mas\u001b[39;00m i16\n\u001b[32m     47\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01m.\u001b[39;00m\u001b[34;01m_binary\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m i32be \u001b[38;5;28;01mas\u001b[39;00m i32\n",
      "\u001b[36mFile \u001b[39m\u001b[32m~/.conda/envs/chromni/lib/python3.11/site-packages/PIL/ImagePalette.py:24\u001b[39m\n\u001b[32m     21\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mcollections\u001b[39;00m\u001b[34;01m.\u001b[39;00m\u001b[34;01mabc\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m Sequence\n\u001b[32m     22\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01mtyping\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m IO\n\u001b[32m---> \u001b[39m\u001b[32m24\u001b[39m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[34;01m.\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m GimpGradientFile, GimpPaletteFile, ImageColor, PaletteFile\n\u001b[32m     26\u001b[39m TYPE_CHECKING = \u001b[38;5;28;01mFalse\u001b[39;00m\n\u001b[32m     27\u001b[39m \u001b[38;5;28;01mif\u001b[39;00m TYPE_CHECKING:\n",
      "\u001b[36mFile \u001b[39m\u001b[32m<frozen importlib._bootstrap>:1176\u001b[39m, in \u001b[36m_find_and_load\u001b[39m\u001b[34m(name, import_)\u001b[39m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m<frozen importlib._bootstrap>:1147\u001b[39m, in \u001b[36m_find_and_load_unlocked\u001b[39m\u001b[34m(name, import_)\u001b[39m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m<frozen importlib._bootstrap>:690\u001b[39m, in \u001b[36m_load_unlocked\u001b[39m\u001b[34m(spec)\u001b[39m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m<frozen importlib._bootstrap_external>:936\u001b[39m, in \u001b[36mexec_module\u001b[39m\u001b[34m(self, module)\u001b[39m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m<frozen importlib._bootstrap_external>:1032\u001b[39m, in \u001b[36mget_code\u001b[39m\u001b[34m(self, fullname)\u001b[39m\n",
      "\u001b[36mFile \u001b[39m\u001b[32m<frozen importlib._bootstrap_external>:1131\u001b[39m, in \u001b[36mget_data\u001b[39m\u001b[34m(self, path)\u001b[39m\n",
      "\u001b[31mKeyboardInterrupt\u001b[39m: "
     ]
    }
   ],
   "source": [
    "import yaml\n",
    "import os \n",
    "import numpy as np\n",
    "\n",
    "from src.ChromnitronDataset import ChromnitronDataset, create_from_original_data\n",
    "from src.utils.vis import plot_histogram, plot_chipseq_tracks, plot_variability \n",
    "import re\n",
    "import argparse\n",
    "\n",
    "config_path = 'config_sene_vs_growing.yaml'\n",
    "config_name=re.sub(r'^(?:config_)?|\\.ya?ml$', '',config_path)\n",
    "# Load config\n",
    "with open(config_path, 'r') as f:\n",
    "    config = yaml.safe_load(f)\n",
    "\n",
    "input_path = config['extract_data_config']['output']['path']\n",
    "output_path_root = config['postprocess_config']['output']['path']\n",
    "\n",
    "data_types = ['upregulated', 'downregulated', 'conserved']\n",
    "\n",
    "print(f\"\\n{'='*60}\")\n",
    "print(f\"Processing {data_type.upper()} data...\")\n",
    "print(f\"{'='*60}\")\n",
    "\n",
    "output_path = os.path.join(output_path_root,config_name, data_type)\n",
    "os.makedirs(output_path, exist_ok=True)\n",
    "npy_path = os.path.join(input_path,config_name, f\"{data_type}_data.npy\")\n",
    "gene_list_path = os.path.join(input_path,config_name, f\"{data_type}_gene_list_no_chrX.csv\")\n",
    "cap_list_path = original_cap_path = os.path.join(config['merge_inference_config']['output']['path'],config['loci_source_config']['atac']['data_path']['experimental'].split(',')[0], 'cap_list.txt')\n",
    "celltypes_list_path = os.path.join(input_path, config_name,f\"{data_type}_celltypes.txt\")\n",
    "\n",
    "### STEP 1: Create ChromnitronDataset\n",
    "print(f\"Loading dataset from {npy_path}...\")\n",
    "dataset = create_from_original_data(npy_path, gene_list_path, cap_list_path, celltypes_list_path, mode='max')\n",
    "\n",
    "\n",
    "data_file=npy_path\n",
    "gene_file=gene_list_path\n",
    "cap_file=cap_list_path\n",
    "celltype_file=celltypes_list_path\n",
    "mode = 'max'\n",
    "\n",
    "# Load data\n",
    "raw_data = np.load(data_file)  # Shape: (n_celltypes, n_loci, n_caps, 2)\n",
    "raw_data = raw_data.reshape(2, -1, *raw_data.shape[1:]).mean(axis=1)  \n",
    "# Load metadata\n",
    "gene_df = pd.read_csv(gene_file)\n",
    "gene_names = gene_df['transcript_name'].tolist()\n",
    "\n",
    "cap_names = pd.read_csv(cap_file, header=None)[0].tolist()\n",
    "\n",
    "with open(celltype_file, 'r') as f:\n",
    "    celltype_names = f.read().strip().split('\\n')\n",
    "\n",
    "# Extract matrices for two cell types\n",
    "if len(celltype_names) < 2:\n",
    "    raise ValueError(\"Need at least 2 cell types\")\n",
    "\n",
    "# Select max or mean values\n",
    "if mode == 'max':\n",
    "    control_matrix = raw_data[1, :, :, 0]     # Second celltype, max values\n",
    "    experiment_matrix = raw_data[0, :, :, 0]  # First celltype, max values\n",
    "    \n",
    "elif mode == 'mean':\n",
    "    control_matrix = raw_data[1, :, :, 1]     # Second celltype, mean values\n",
    "    experiment_matrix = raw_data[0, :, :, 1]  # First celltype, mean values\n",
    "else:\n",
    "    raise ValueError(\"Mode must be 'max' or 'mean'\")\n",
    "\n",
    "# Create dataset\n",
    "dataset = ChromnitronDataset(control_matrix, experiment_matrix, gene_names, cap_names,\n",
    "                             config_name.split('_vs_')[0],config_name.split('_vs_')[1]) #hua changed\n",
    "                            #celltype_names[1], celltype_names[0])\n",
    "\n",
    "\n",
    "\n",
    "\n",
    "### STEP 2: Compute differential matrix\n",
    "print(\"Computing log2 fold change differential...\")\n",
    "dataset.compute_differential(log_transform=True, base='2')\n",
    "\n",
    "# Histogram (if enabled)\n",
    "if config['postprocess_config']['plotting'].get('plot_histograms', True):\n",
    "    print(\"Generating histogram after differential computation...\")\n",
    "    diff_matrix = dataset.get_differential()\n",
    "    plot_histogram(\n",
    "        data=diff_matrix.flatten(),\n",
    "        bins=100,\n",
    "        plot_title=f\"{data_type.capitalize()} - Log2 Fold Change Distribution\",\n",
    "        x_name=\"Log2 Fold Change\",\n",
    "        y_name=\"Frequency\",\n",
    "        filename=f\"{data_type}_histogram_1_differential\",\n",
    "        data_save_path=output_path\n",
    "    )\n",
    "    print(f\"Saved differential histogram to {output_path}/{data_type}_histogram_1_differential.png\")\n",
    "else:\n",
    "    print(\"Skipping histogram generation (disabled in config)\")\n",
    "\n",
    "### STEP 3: Variance filtering (if enabled)\n",
    "if config['postprocess_config']['variance_filtering'].get('enable', True):\n",
    "    n_genes = config['postprocess_config']['variance_filtering'].get('n_genes', 500)\n",
    "    n_caps = config['postprocess_config']['variance_filtering'].get('n_caps', 100)\n",
    "    print(f\"Applying variance filtering: top {n_genes} genes, top {n_caps} CAPs...\")\n",
    "    var_results = dataset.variance_filter(n_genes=n_genes, n_caps=n_caps, condition='differential')\n",
    "else:\n",
    "    print(\"Skipping variance filtering (disabled in config)\")\n",
    "    # Create a mock var_results for plotting if variability plot is enabled\n",
    "    var_results = None\n",
    "\n",
    "# Histogram (if both variance filtering and histograms are enabled)\n",
    "if (config['postprocess_config']['variance_filtering'].get('enable', True) and \n",
    "    config['postprocess_config']['plotting'].get('plot_histograms', True)):\n",
    "    print(\"Generating histogram after variance filtering...\")\n",
    "    filtered_diff_matrix = dataset.get_differential()\n",
    "    n_genes = config['postprocess_config']['variance_filtering'].get('n_genes', 500)\n",
    "    n_caps = config['postprocess_config']['variance_filtering'].get('n_caps', 100)\n",
    "    plot_histogram(\n",
    "        data=filtered_diff_matrix.flatten(),\n",
    "        bins=100,\n",
    "        plot_title=f\"{data_type.capitalize()} - After Variance Filtering (Top {n_genes} genes, {n_caps} CAPs)\",\n",
    "        x_name=\"Log2 Fold Change\",\n",
    "        y_name=\"Frequency\",\n",
    "        filename=f\"{data_type}_histogram_2_variance_filtered\",\n",
    "        data_save_path=output_path\n",
    "    )\n",
    "    print(f\"Saved variance filtered histogram to {output_path}/{data_type}_histogram_2_variance_filtered.png\")\n",
    "else:\n",
    "    if not config['postprocess_config']['variance_filtering'].get('enable', True):\n",
    "        print(\"Skipping variance filtering histogram (variance filtering disabled)\")\n",
    "    else:\n",
    "        print(\"Skipping variance filtering histogram (plotting disabled in config)\")\n",
    "\n",
    "# Plot variability (if both variance filtering and variability plotting are enabled)\n",
    "if (config['postprocess_config']['variance_filtering'].get('enable', True) and \n",
    "    config['postprocess_config']['plotting'].get('plot_variability', True) and \n",
    "    var_results is not None):\n",
    "    variability_plot_path = os.path.join(output_path, f\"{data_type}_variability.png\")\n",
    "    plot_variability(var_results, save_path=variability_plot_path, show_labels=True, max_labels=50)\n",
    "    print(f\"Saved variability plot to {variability_plot_path}\")\n",
    "else:\n",
    "    if not config['postprocess_config']['variance_filtering'].get('enable', True):\n",
    "        print(\"Skipping variability plot (variance filtering disabled)\")\n",
    "    else:\n",
    "        print(\"Skipping variability plot (plotting disabled in config)\")\n",
    "\n",
    "### Plot ChIP-seq tracks before noise removal (if enabled)\n",
    "if config['postprocess_config']['plotting'].get('plot_chipseq', True):\n",
    "    print(\"Generating ChIP-seq tracks before noise removal...\")\n",
    "    gene_names = dataset.filtered_gene_names or dataset.gene_names\n",
    "    cap_names = dataset.filtered_cap_names or dataset.cap_names\n",
    "    \n",
    "    # Plot all three conditions\n",
    "    conditions_to_plot = [\n",
    "        (dataset.get_control(), dataset.control_name, \"raw\"),\n",
    "        (dataset.get_experiment(), dataset.experiment_name, \"raw\"), \n",
    "        (dataset.get_differential(), \"Differential\", \"log2fc\")\n",
    "    ]\n",
    "    \n",
    "    for matrix, condition_name, data_type_suffix in conditions_to_plot:\n",
    "        plot_chipseq_tracks(\n",
    "            matrix, \n",
    "            cap_names,\n",
    "            loci_list=gene_names,\n",
    "            n_caps_per_page=100,  # Smaller pages for multiple conditions\n",
    "            track_height=0.25,   # Smaller tracks\n",
    "            figsize_width=250,   # Smaller width\n",
    "            smooth_window=3,     # Less smoothing for raw data\n",
    "            log_scale=False,\n",
    "            show_percentiles=True,\n",
    "            data_save_path=output_path,\n",
    "            filename_prefix=f\"{data_type}_{condition_name.lower().replace(' ', '_')}_before_noise_removal\",\n",
    "            output_format=\"pdf\"\n",
    "        )\n",
    "        print(f\"Saved {condition_name} tracks (before noise removal) to {output_path}\")\n",
    "else:\n",
    "    print(\"Skipping ChIP-seq tracks before noise removal (disabled in config)\")\n",
    "\n",
    "### STEP 4: Remove noise by keeping only top X% of signal (if enabled)\n",
    "if config['postprocess_config']['remove_noise'].get('enable', True):\n",
    "    top_percent = config['postprocess_config']['remove_noise'].get('top_percent', 10.0)\n",
    "    conditions = config['postprocess_config']['remove_noise'].get('conditions', ['all'])\n",
    "    print(f\"Removing noise (keeping top {top_percent}% extreme values)...\")\n",
    "    dataset.remove_noise(top_percent=top_percent, conditions=conditions)\n",
    "else:\n",
    "    print(\"Skipping noise removal (disabled in config)\")\n",
    "\n",
    "# Histogram (if both noise removal and histograms are enabled)\n",
    "if (config['postprocess_config']['remove_noise'].get('enable', True) and \n",
    "    config['postprocess_config']['plotting'].get('plot_histograms', True)):\n",
    "    print(\"Generating histogram after noise removal...\")\n",
    "    final_diff_matrix = dataset.get_differential()\n",
    "    top_percent = config['postprocess_config']['remove_noise'].get('top_percent', 10.0)\n",
    "    plot_histogram(\n",
    "        data=final_diff_matrix.flatten(),\n",
    "        bins=100,\n",
    "        plot_title=f\"{data_type.capitalize()} - After Noise Removal (Top {top_percent}% Extreme Values)\",\n",
    "        x_name=\"Log2 Fold Change\",\n",
    "        y_name=\"Frequency\",\n",
    "        filename=f\"{data_type}_histogram_3_noise_removed\",\n",
    "        data_save_path=output_path\n",
    "    )\n",
    "    print(f\"Saved noise removed histogram to {output_path}/{data_type}_histogram_3_noise_removed.png\")\n",
    "else:\n",
    "    if not config['postprocess_config']['remove_noise'].get('enable', True):\n",
    "        print(\"Skipping noise removal histogram (noise removal disabled)\")\n",
    "    else:\n",
    "        print(\"Skipping noise removal histogram (plotting disabled in config)\")\n",
    "\n",
    "### Plot ChIP-seq tracks AFTER noise removal (if both noise removal and chipseq plotting are enabled)\n",
    "if (config['postprocess_config']['remove_noise'].get('enable', True) and \n",
    "    config['postprocess_config']['plotting'].get('plot_chipseq', True)):\n",
    "    print(\"Generating ChIP-seq tracks AFTER noise removal...\")\n",
    "    \n",
    "    # Plot all three conditions after noise removal\n",
    "    conditions_to_plot_after = [\n",
    "        (dataset.get_control(), dataset.control_name, \"raw_clean\"),\n",
    "        (dataset.get_experiment(), dataset.experiment_name, \"raw_clean\"), \n",
    "        (dataset.get_differential(), \"Differential\", \"log2fc_clean\")\n",
    "    ]\n",
    "    \n",
    "    for matrix, condition_name, data_type_suffix in conditions_to_plot_after:\n",
    "        plot_chipseq_tracks(\n",
    "            matrix, \n",
    "            cap_names,\n",
    "            loci_list=gene_names,\n",
    "            n_caps_per_page=100,  # Smaller pages for multiple conditions\n",
    "            track_height=0.25,   # Smaller tracks\n",
    "            figsize_width=250,   # Smaller width\n",
    "            smooth_window=3,     # Less smoothing\n",
    "            log_scale=False,\n",
    "            show_percentiles=True,\n",
    "            data_save_path=output_path,\n",
    "            filename_prefix=f\"{data_type}_{condition_name.lower().replace(' ', '_')}_after_noise_removal\",\n",
    "            output_format=\"pdf\"\n",
    "        )\n",
    "        print(f\"Saved {condition_name} tracks (after noise removal) to {output_path}\")\n",
    "else:\n",
    "    if not config['postprocess_config']['remove_noise'].get('enable', True):\n",
    "        print(\"Skipping ChIP-seq tracks after noise removal (noise removal disabled)\")\n",
    "    else:\n",
    "        print(\"Skipping ChIP-seq tracks after noise removal (plotting disabled in config)\")\n",
    "\n",
    "### STEP 5: Save processed dataset\n",
    "print(\"Saving processed dataset...\")\n",
    "dataset.save(output_path)\n",
    "print(f\"Saved dataset to {output_path}\")\n",
    "\n",
    "# Summary\n",
    "print(\"Dataset summary:\")\n",
    "dataset.summary()\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a603d631-04bb-475f-9921-9bfd2dd59c7a",
   "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
}
