Skip to content
Merged
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
25 changes: 24 additions & 1 deletion config/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,10 @@ splicing_consequences:
models_experiment1:
- CADD.plus.score
- GPN-MSA_LLR.minus.score
- PromoterAI.plus.score
- GPN-Star-V_LLR.minus.llr_calibrated
- GPN-Star-M_LLR.minus.llr_calibrated
- GPN-Star-P_LLR.minus.llr_calibrated
# - PromoterAI.plus.score

datasets:
- mendelian_traits_matched_9
Expand Down Expand Up @@ -374,3 +377,23 @@ evo2:
model_path: evo2_40b
window_size: 8192
per_device_batch_size: 8 # H200


gpn_star:
V:
repo_id: songlab/gpn-star-hg38-v100-200m
msa_path: /scratch/users/czye/GPN/egpn/analysis/human/egpn/workflow/results/msa/hg38/multiz100way/100
window_size: 128
per_device_batch_size: 64

M:
repo_id: songlab/gpn-star-hg38-m447-200m
msa_path: /scratch/users/czye/GPN/egpn/analysis/human/egpn/workflow/results/msa/hg38/cactus447way/447
window_size: 256
per_device_batch_size: 64

P:
repo_id: songlab/gpn-star-hg38-p243-200m
msa_path: /scratch/users/czye/GPN/egpn/analysis/human/egpn/workflow/results/msa/hg38/cactus447way/243
window_size: 256
per_device_batch_size: 64
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,9 @@ dependencies = [
"scipy",
"seaborn",
"scikit-learn",
"huggingface_hub",
"huggingface-hub",
"liftover",
"gpn @ git+https://github.com/songlab-cal/gpn.git",
"gpn @ git+https://github.com/songlab-cal/gpn.git@fallback-phylo-dist-path",
"cyvcf2>=0.31.4",
"openpyxl>=3.1.5",
"gnomad-db @ git+https://github.com/KalinNonchev/gnomAD_DB.git",
Expand Down
4 changes: 2 additions & 2 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

1 change: 1 addition & 0 deletions workflow/Snakefile
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ include: "rules/features/enformer.smk"
include: "rules/features/evo2.smk"
include: "rules/features/gpn.smk"
include: "rules/features/gpn_msa.smk"
include: "rules/features/gpn_star.smk"
include: "rules/features/grelu.smk"
include: "rules/features/hyenadna.smk"
include: "rules/features/maf.smk"
Expand Down
7 changes: 7 additions & 0 deletions workflow/rules/data/mendelian_traits.smk
Original file line number Diff line number Diff line change
Expand Up @@ -142,3 +142,10 @@ rule mendelian_traits_all_dataset:
run:
V = pl.read_parquet(input[0])
V.sort(COORDINATES).write_parquet(output[0])


rule mendelian_traits_legacy_dataset:
output:
"results/dataset/mendelian_traits_legacy/test.parquet",
shell:
"wget -O {output} https://huggingface.co/datasets/songlab/TraitGym/resolve/main/mendelian_traits_matched_9/test.parquet"
66 changes: 66 additions & 0 deletions workflow/rules/features/gpn_star.smk
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
rule gpn_star_download_model:
output:
directory("results/gpn_star/checkpoints/{model}"),
params:
repo_id=lambda wildcards: config["gpn_star"][wildcards.model]["repo_id"],
threads: workflow.cores
shell:
"hf download {params.repo_id} --local-dir {output} --max-workers {threads}"


# intermediate output used to obtain both LLR, entropy
rule get_logits:
input:
"results/dataset/{dataset}/test.parquet",
lambda wildcards: config["gpn_star"][wildcards.model]["msa_path"],
"results/gpn_star/checkpoints/{model}",
output:
"results/gpn_star/logits/{dataset}/{model}.parquet",
params:
window_size=lambda wildcards: config["gpn_star"][wildcards.model]["window_size"],
resources:
# using resources to avoid re-runs when changing batch size
per_device_batch_size=lambda wildcards: config["gpn_star"][wildcards.model][
"per_device_batch_size"
],
threads: workflow.cores
shell:
"""
python \
-m gpn.star.inference logits \
{input[0]} {input[1]} {params.window_size} {input[2]} {output} \
--per_device_batch_size {resources.per_device_batch_size} \
--dataloader_num_workers {threads} \
--is_file
"""


rule get_llr_calibrated:
input:
"results/dataset/{dataset}/test.parquet",
"results/genome.fa.gz",
"results/gpn_star/checkpoints/{model}",
"results/gpn_star/logits/{dataset}/{model}.parquet",
output:
"results/dataset/{dataset}/features/GPN-Star-{model}_LLR.parquet",
params:
calibration_path="results/gpn_star/checkpoints/{model}/calibration_table/llr.parquet",
run:
from gpn.star.utils import normalize_logits, get_llr

V = pd.read_parquet(input[0])
genome = Genome(input[1])
V["pentanuc"] = V.apply(
lambda row: genome.get_seq(
row["chrom"], row["pos"] - 3, row["pos"] + 2
).upper(),
axis=1,
)
V["pentanuc_mut"] = V["pentanuc"] + "_" + V["alt"]
df_calibration = pd.read_parquet(params.calibration_path)
logits = pd.read_parquet(input[3])
normalized_logits = normalize_logits(logits)
V["llr"] = get_llr(normalized_logits, V["ref"], V["alt"])
V = V.merge(df_calibration, on="pentanuc_mut", how="left")
V["llr_calibrated"] = V["llr"] - V["llr_neutral_mean"]
V[["llr", "llr_calibrated"]].to_parquet(output[0], index=False)