Skip to content

Commit d9360d8

Browse files
committed
feat: filter to top-N trees by llh in generate_input main loop
1 parent e57a45d commit d9360d8

2 files changed

Lines changed: 28 additions & 1 deletion

File tree

src/neoantigen_utils/generate_input.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -347,7 +347,7 @@ def main(args):
347347

348348
trees = summ_data["trees"]
349349

350-
for tree in trees:
350+
for tree in select_top_trees(trees, args.top_n_trees):
351351

352352
inner_sample_tree_dict = {"topology": [], "score": trees[tree]["llh"]}
353353
with open("./" + args.tree_directory + "/" + str(tree) + ".json", "r") as f:
@@ -1331,6 +1331,12 @@ def parse_args():
13311331
parser.add_argument(
13321332
"--kD_cutoff", default=500, help="Cutoff value for the kD, default is 500",
13331333
)
1334+
parser.add_argument(
1335+
"--top_n_trees",
1336+
type=int,
1337+
default=10,
1338+
help="Number of top trees (by llh, highest first) to include, default is 10",
1339+
)
13341340

13351341
return parser.parse_args()
13361342

tests/test_generate_input.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -230,3 +230,24 @@ def test_n_zero_returns_empty_list(self):
230230

231231
def test_empty_trees_returns_empty_list(self):
232232
assert select_top_trees({}, 10) == []
233+
234+
235+
class TestTopNTreesArgDefault:
236+
def test_default_top_n_trees_is_ten(self):
237+
import sys
238+
from neoantigen_utils.generate_input import parse_args
239+
240+
argv = [
241+
"prog",
242+
"--maf_file", "x", "--summary_file", "x", "--mutation_file", "x",
243+
"--tree_directory", "x", "--id", "x", "--patient_id", "x",
244+
"--cohort", "x", "--HLA_genes", "x",
245+
"--netMHCpan_MUT_input", "x", "--netMHCpan_WT_input", "x",
246+
]
247+
old_argv = sys.argv
248+
sys.argv = argv
249+
try:
250+
args = parse_args()
251+
finally:
252+
sys.argv = old_argv
253+
assert args.top_n_trees == 10

0 commit comments

Comments
 (0)