Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 59 additions & 0 deletions 2DTM_postprocess_tool/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
# 2DTM Postprocessing

A modular Python package for postprocessing 2D template matching results from cryo-EM workflows (e.g., cisTEM), including 2DTM p-value calculation, particle extraction and filtering.

---

## Installation

```bash
git clone https://github.com/kekexinz/2DTM_postprocess_tool.git
cd 2DTM_postprocess_tool
pip install -e . # editable mode
```

## 📦 Usage

### `extract-particles`
Extract initial particle peaks from 2DTM search.
```bash
extract-particles \
--db_file <cistem.db> \
--tm_job_id 1 \
--ctf_job_id 1 \
--pixel_size 1.0 \
--output <extracted_peaks.star>
[--metric pval] \ # "zscore" or "pval"
[--metric_cutoff 8.0] \
[--threads 22] \
[--local_max_filter] \ # "snr" or "zscore" (default) used for skimage peak_local_max
[--min_peak_radius 10] \ # used for "min_distance" in skimage peak_local_max
[--exclude_borders 92] \ # avoid finding partial particles near the edge of the image, used for skimage peak_local_max
[--quadrants 1] \ # 1 (default) or 3, calculating p-value for only the first-quadrant or quadrant 1,2,4 (recommended for small particles)

```

### `filter-particles`

Filter particles based on image thickness and/or angular invariance.

```bash
filter-particles \
--star_file <extracted_peaks.star> \ # output from extract-particles
--db_file <cistem.db> \
--tm_job_id 1 \
--ctf_job_id 1 \
--pixel_size 1.0 \
--output filtered_peaks.star \
[--avg_cutoff_lb] \ # angular search CC per-pixel avg
[--sd_cutoff_ub] \ # angular search CC per-pixel sd
[--snr_cutoff_ub] \
[--filter_by_image_thickness] \ # ctffind5 parameters
[--thickness_cutoff_lb] \
[--thickness_cutoff_ub] \
[--ctf_fitting_score_lb] \
[--ctf_fitting_score_ub] \
```

### 3D reconstruction & refinement in cisTEM
The output extracted_peaks.star and filtered_peaks.star can be imported into cisTEM as a RefinementPackage for further 3D reconstruction and refinement.
25 changes: 25 additions & 0 deletions 2DTM_postprocess_tool/setup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
from setuptools import setup, find_packages

setup(
name="tm_post",
version="0.1",
author="Kexin Zhang",
description="Postprocessing utilities for 2D template matching",
packages=find_packages(where="src"),
package_dir={"": "src"},
install_requires=[
"numpy",
"pandas",
"scipy",
"mrcfile",
"joblib",
"tqdm",
"scikit-image",
],
entry_points={
"console_scripts": [
"filter-particles = cli.filter_particles:main",
"extract-particles = cli.extract_particles:main",
],
},
)
Empty file.
42 changes: 42 additions & 0 deletions 2DTM_postprocess_tool/src/cli/compare_starfiles.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
import argparse
import pandas as pd
import tm_post.starfile as starfile
from tm_post.compare_starfiles import compare_starfiles_for_matched_peaks

def parse_arguments():
parser = argparse.ArgumentParser(
description="Compare two STAR files and extract matched peaks based on spatial and angular thresholds."
)
parser.add_argument('--starfile_a', type=str, required=True, help="Path to first STAR file (e.g. bin2x).")
parser.add_argument('--starfile_b', type=str, required=True, help="Path to second STAR file (e.g. bin1x).")
parser.add_argument('--d_xy_cutoff', type=float, default=10.0, help="Maximum XY distance in Å for matching.")
parser.add_argument('--euler_err_cutoff', type=float, default=5.0, help="Maximum Euler angle error in degrees.")
parser.add_argument('--pattern', type=str, default=r"mc2_[12]x_(.*?frames)", help="Regex pattern for extracting match key.")
parser.add_argument('--output', type=str, required=True, help="Path to output STAR file with matched peaks.")

return parser.parse_args()

def main():
args = parse_arguments()

print("[INFO] Comparing starfiles...")
matched_df = compare_starfiles_for_matched_peaks(
starfile_a=args.starfile_a,
starfile_b=args.starfile_b,
d_xy_cutoff=args.d_xy_cutoff,
euler_err_cutoff=args.euler_err_cutoff,
pattern=args.pattern
)

if matched_df.empty:
print("[INFO] No matching peaks found.")
return

# Convert matched_df to STAR format (with dummy column) and write with standard header
matched_df_star = starfile.add_star_dummy_column(matched_df)
header_lines = starfile.read_tm_package_starfile_header()
starfile.write_starfile_with_headers(args.output, header_lines, matched_df_star)
print(f"[INFO] Matched particles saved to: {args.output}")

if __name__ == "__main__":
main()
65 changes: 65 additions & 0 deletions 2DTM_postprocess_tool/src/cli/extract_particles.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
import argparse
import pandas as pd
import tm_post.database as db
from tm_post.database import load_tm_images_from_db
from tm_post.extract import extract_particles_from_2dtm_search
from tm_post.starfile import write_starfile_with_headers, read_tm_package_starfile_header

def parse_arguments():
parser = argparse.ArgumentParser(description="Extract peaks from 2DTM searches.")

parser.add_argument('--db_file', type=str, required=True, help="Path to the database file.")
parser.add_argument('--tm_job_id', type=int, required=True, help="Template match job ID.")
parser.add_argument('--ctf_job_id', type=int, required=True, help="CTF job ID.")

parser.add_argument('--min_peak_radius', type=int, default=10, help="Cutoff for XY distance.")
parser.add_argument('--exclude_borders', type=int, default=35, help="Exclude borders in the image.")

parser.add_argument('--local_max_filter', type=str, default="zscore", choices=["zscore", "snr"], help="Local max filter to use.")
parser.add_argument('--metric', type=str, default="pval", choices=["pval", "zscore", "snr"], help="Metric to use for filtering.")
parser.add_argument('--metric_cutoff', type=float, default=8.0, help="Selected metric cutoff.")
parser.add_argument('--pixel_size', type=float, required=True, default=1.0, help="Wanted pixel size in final stack.")
parser.add_argument('--threads', type=int, default=4, help="Number of threads for parallel processing.")

parser.add_argument('--quadrants', type=int, default=1, help="Number of quadrants to use for filtering.")

parser.add_argument('--output', type=str, required=True, help="Path to the output star file.")

return parser.parse_args()

def main():
args = parse_arguments()

# Load database information
print("[INFO] Loading TM image data from database...")
tm_images, df_ctf, df_info = load_tm_images_from_db(
db_file=args.db_file,
tm_job_id=args.tm_job_id,
ctf_job_id=args.ctf_job_id
)

print(f"[INFO] Running extraction on {len(tm_images)} images...")

df_star = extract_particles_from_2dtm_search(
tm_images=tm_images,
local_max_filter=args.local_max_filter,
metric=args.metric,
metric_cutoff=args.metric_cutoff,
pixel_size=args.pixel_size,
max_threads=args.threads,
df_ctf=df_ctf,
df_info=df_info,
ctf_job_id=args.ctf_job_id,
min_radius=args.min_peak_radius,
exclude_borders=args.exclude_borders,
q=args.quadrants
)

print("[INFO] Writing STAR file...")
header_lines = read_tm_package_starfile_header() # provide default STAR header
write_starfile_with_headers(args.output, header_lines, df_star)
print(f"[INFO] Done. Extracted particles saved to {args.output}")


if __name__ == "__main__":
main()
121 changes: 121 additions & 0 deletions 2DTM_postprocess_tool/src/cli/filter_particles.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,121 @@
import argparse
import tm_post.starfile as starfile
import tm_post.database as db
from tm_post.filters import apply_filter

def parse_arguments():
parser = argparse.ArgumentParser(description="Filter peaks using TM results and image quality.")
# read job information from database file
parser.add_argument('--star_file', type=str, required=True, help="Path to the particle starfile.")
parser.add_argument('--db_file', type=str, required=True, help="Path to the database file.")
parser.add_argument('--tm_job_id', type=int, required=True, help="Template match job ID.")
parser.add_argument('--ctf_job_id', type=int, required=True, help="CTF job ID.")
parser.add_argument('--pixel_size', type=float, required=True, help="Pixel size in Angstroms.")

# read particle information from .star file (extract_peaks.py output)
parser.add_argument('--avg_cutoff_lb', type=float, default=None, help="Lower bound for average cutoff.")
parser.add_argument('--sd_cutoff_ub', type=float, default=None, help="Upper bound for SD cutoff.")
parser.add_argument('--pval_cutoff_lb', type=float, default=None, help="Lower bound for p-value cutoff.")
parser.add_argument('--snr_cutoff_ub', type=float, default=None, help="Upper bound for SNR (optional).")
parser.add_argument('--snr_cutoff_lb', type=float, default=None, help="Lower bound for SNR (optional).")
parser.add_argument('--filter_by_image_thickness', action="store_true", help="Use thickness to filter good micrographs? (default: False)")
parser.add_argument('--thickness_cutoff_lb', type=float, default=None, help="Lower bound for thickness cutoff (A).")
parser.add_argument('--thickness_cutoff_ub', type=float, default=None, help="Upper bound for thickness cutoff (A).")
parser.add_argument('--filter_by_angular_invariance', action="store_true", help="Use angular invariance to filter good particles (default: False)?")
parser.add_argument('--geodesic_r', type=int, default=None, help="Radius in pixels for local patch.")
parser.add_argument('--geodesic_threads', type=int, default=None, help="Number of threads for geodesic computation.")
parser.add_argument('--geodesic_method', type=str, default=None, help="Method for geodesic filtering ('quantile' or 'cutoff').")
parser.add_argument('--geodesic_threshold', type=float, default=None, help="Threshold value for geodesic filtering.")

parser.add_argument('--ctf_fitting_score_lb', type=float, default=None, help="Lower bound for CTF fitting score.")
parser.add_argument('--ctf_fitting_score_ub', type=float, default=None, help="Upper bound for CTF fitting score.")


parser.add_argument('--output', type=str, required=True, help="Path to the output star file.")
return parser.parse_args()


def main():
args = parse_arguments()

# Load particle information from star file
print("[INFO] Loading peak file...")
df_peaks = starfile.load_particle_starfile(args.star_file)

# Extract header lines
#header_lines, _ = starfile.extract_header_lines(args.star_file)

# Load database information
print("[INFO] Loading database...")
result = db.get_info_from_cistem_database(
args.db_file, args.tm_job_id, args.ctf_job_id
)

# Extract relevant data
image_list = result["image_list"]
psi_list = result["PSI_OUTPUT_FILE"]
theta_list = result["THETA_OUTPUT_FILE"]
phi_list = result["PHI_OUTPUT_FILE"]
df_ctf = result["df_ctf"]
df_info = result["df_info"]

# Apply filters
if args.filter_by_image_thickness:
if args.thickness_cutoff_lb is None:
args.thickness_cutoff_lb = 0.0
if args.thickness_cutoff_ub is None:
args.thickness_cutoff_ub = 500.0

if args.filter_by_angular_invariance:
if args.geodesic_r is None:
args.geodesic_r = 4
if args.geodesic_threads is None:
args.geodesic_threads = 8
if args.geodesic_method is None:
args.geodesic_method = "quantile"
if args.geodesic_threshold is None:
args.geodesic_threshold = 0.8

filtered_df, all_df = apply_filter(
df=df_peaks,
image_list=image_list,
psi_list=psi_list,
theta_list=theta_list,
phi_list=phi_list,
pixel_size=args.pixel_size,
df_ctf=df_ctf,
df_info=df_info,
avg_cutoff_lb=args.avg_cutoff_lb,
sd_cutoff_ub=args.sd_cutoff_ub,
pval_cutoff_lb=args.pval_cutoff_lb,
snr_cutoff_ub=args.snr_cutoff_ub,
snr_cutoff_lb=args.snr_cutoff_lb,
filter_by_image_thickness=args.filter_by_image_thickness,
thickness_lb=args.thickness_cutoff_lb,
thickness_ub=args.thickness_cutoff_ub,
ctf_fitting_score_lb=args.ctf_fitting_score_lb,
ctf_fitting_score_ub=args.ctf_fitting_score_ub,
filter_by_angular_invariance=args.filter_by_angular_invariance,
geodesic_r=args.geodesic_r,
geodesic_threads=args.geodesic_threads,
geodesic_method=args.geodesic_method,
geodesic_threshold=args.geodesic_threshold
)

# Save filtered STAR file with updated SCORE (no extra metadata columns)
columns_to_keep = ["ORIGINAL_IMAGE_FILENAME","ORIGX","ORIGY","AVG","SD","PVALUE","ZSCORE","SNR"]
meta_df = all_df[columns_to_keep].copy()

# Convert filtered DataFrame to lines with empty column
starfile_df_star = starfile.add_star_dummy_column(filtered_df)

# Write output with headers
header_lines = starfile.read_tm_package_starfile_header() # provide default STAR header
starfile.write_starfile_with_headers(args.output, header_lines, starfile_df_star)
print(f"[INFO] Filtered data saved to {args.output}")

metadata_file = args.output.replace(".star", "_metadata.txt")
meta_df.to_csv(metadata_file, sep="\t", index=False, float_format="%.2f")

if __name__ == "__main__":
main()
70 changes: 70 additions & 0 deletions 2DTM_postprocess_tool/src/cli/update_par.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
#!/usr/bin/env python3

import argparse
import pandas as pd
from io import StringIO
from tm_post import starfile

def read_par_file(par_path):
"""Read a cisTEM .par file into header, data DataFrame, and footer."""
with open(par_path, "r") as f:
lines = f.readlines()

header = lines[0]
footer = lines[-2:]
data_lines = lines[1:-2]

data_str = ''.join(data_lines)
df = pd.read_csv(StringIO(data_str), delim_whitespace=True, header=None)
df.columns = header.strip().split()

return header, df, footer


def read_score_file(score_path):
"""Read a file with one score per line."""
df = starfile.load_particle_starfile(score_path)
return df["SCORE"].tolist()


def update_scores(df, score_files):
"""Concatenate scores from all files and assign to the SCORE column."""
all_scores = []
for path in score_files:
scores = read_score_file(path)
all_scores.extend(scores)

if len(all_scores) != len(df):
raise ValueError(f"Number of scores ({len(all_scores)}) does not match number of particles ({len(df)}).")

df['SCORE'] = all_scores
return df


def write_par_file(out_path, header, df, footer):
"""Write the updated .par file."""
with open(out_path, "w") as f:
f.write(header)
for row in df.itertuples(index=False):
values = ' '.join(f"{v:>8}" if isinstance(v, float) else f"{v:>8}" for v in row)
f.write(f"{values}\n")
f.writelines(footer)


def main():
parser = argparse.ArgumentParser(description="Update SCORE column in a cisTEM .par file.")
parser.add_argument("par_file", help="Path to the original .par file")
parser.add_argument("score_files", nargs='+', help="One or more star files containing updated SCORE values")
parser.add_argument("-o", "--output", required=True, help="Output path for updated .par file")

args = parser.parse_args()

header, df, footer = read_par_file(args.par_file)
df = update_scores(df, args.score_files)
write_par_file(args.output, header, df, footer)

print(f"[INFO] Updated .par file written to: {args.output}")


if __name__ == "__main__":
main()
Empty file.
Loading
Loading