Added custom models as submodule, added pipleline for trap species handling (up to alignment). Initial progress on custom model rules

This commit is contained in:
Tom Kasper
2026-10-03 15:16:36 +01:00
parent 3b93ceaf6e
commit f0cd594148
17 changed files with 270 additions and 21 deletions
+32 -10
View File
@@ -22,9 +22,21 @@ def get_batch_names(input_file):
with open(input_file,'rt') as ih:
return [line.strip().split('/')[-1].split('.')[0] for line in ih]
rule download_pod5:
def get_run_name(wildcards):
if wildcards.dataset == 'nomiss':
return 'PBK98658_853a956f_57f83f46'
if wildcards.dataset == 'trap':
return 'ATCC_25922_202309'
def get_batch_range(wildcards):
if wildcards.dataset == 'nomiss':
return range(1,config['pod5_dataset_size']+1,config['pod5_stride'])
if wildcards.dataset == 'trap':
return [0]
rule download_nomiss_pod5:
output:
'../data/raw_pod5/PBK98658_853a956f_57f83f46_{batch}.pod5'
'../data/raw_pod5/nomiss/PBK98658_853a956f_57f83f46_{batch}.pod5'
threads: 1
wildcard_constraints:
batch="\d+"
@@ -33,12 +45,22 @@ rule download_pod5:
aws s3 cp --no-sign-request s3://ont-open-data/nomiss_96BC_P2I_SUP_2026/raw/pod5/PBK98658_853a956f_57f83f46_{wildcards.batch}.pod5 {output}
"""
rule download_trap_pod5:
output:
'../data/raw_pod5/trap/ATCC_25922_202309_0.pod5'
threads: 1
wildcard_constraints:
batch='\d+'
shell:
"""
curl -L "https://api.figshare.com/v2/file/download/45408628" -o {output}
"""
rule basecall_pod5:
input:
pod5='../data/raw_pod5/PBK98658_853a956f_57f83f46_{batch}.pod5',
pod5='../data/raw_pod5/{dataset}/{run}_{batch}.pod5',
model=get_model_requirement
output:
temp('../data/basecalled_reads/{model}/PBK98658_853a956f_57f83f46_{batch}.fastq')
temp('../data/basecalled_reads/{dataset}/{model}/{run}_{batch}.fastq')
threads:
32
resources:
@@ -57,9 +79,9 @@ rule basecall_pod5:
rule concatenate_basecalled_fastq:
input:
expand('../data/basecalled_reads/{{model}}/PBK98658_853a956f_57f83f46_{batch}.fastq',batch=range(1,config["pod5_dataset_size"]+1,config["pod5_stride"]))
expand('../data/basecalled_reads/{{dataset}}/{{model}}/{run}_{batch}.fastq',batch=get_batch_range,run=get_run_name)
output:
'../data/basecalled_reads/{model}.fastq.gz'
'../data/basecalled_reads/{dataset}_{model}.fastq.gz'
threads: 1
wildcard_constraints:
model="hac|fast"
@@ -70,10 +92,10 @@ rule concatenate_basecalled_fastq:
rule align_reads_to_reference:
input:
fastq='../data/basecalled_reads/{model}.fastq.gz',
fastq='../data/basecalled_reads/{dataset}_{model}.fastq.gz',
ref='../data/reference_genomes/full_reference.mmi'
output:
'../data/aligned_reads/{model}_to_genome.sam'
'../data/aligned_reads/{dataset}_{model}_to_genome.sam'
threads: 32
conda:
'../envs/minimap.yaml'
@@ -84,9 +106,9 @@ rule align_reads_to_reference:
rule convert_sam_to_bam:
input:
'../data/aligned_reads/{model}_to_genome.sam'
'../data/aligned_reads/{dataset}_{model}_to_genome.sam'
output:
'../data/aligned_reads/{model}_to_genome.sorted.bam'
'../data/aligned_reads/{dataset}_{model}_to_genome.sorted.bam'
threads: 32
conda:
'../envs/samtools.yaml'
+66
View File
@@ -0,0 +1,66 @@
rule prepare_read_labels:
input:
bamfile="../data/aligned_reads/{model}_to_genome.sorted.bam",
reference="../data/reference_genomes/full_reference.fasta"
output:
'../data/ml_inputs/{model}_label_store.pq'
conda:
'../envs/bam2parquet.yaml'
threads:
1
script:
'../scripts/bam2annotation.py'
rule build_data_store:
input:
pod5=expand('../data/raw_pod5/{{model}}/PBK98658_853a956f_57f83f46_{batch}.pod5',batch=range(1,config["pod5_dataset_size"]+1,config["pod5_stride"])),
labels='../data/ml_inputs/{model}_label_store.pq'
output:
directory('../data/ml_inputs/{model}_data_store')
conda:
config["torch_env"]
shell:
"""
python -m custom_models store\
--out {output}\
--pod5 {input.pod5}\
--labels {input.labels}\
--reads-per-genus {config[reads_per_genus]}
"""
rule train_model:
input:
data_store='../data/ml_inputs/{model}_data_store'
output:
directory('../data/ml_models/run_{arch}_{params}_{stride}_{embed}_{heads}')
conda:
config["torch_env"]
shell:
"""
python -m custom_models train\
--arch {wildcards.arch}\
--budget {wildcards.params}\
--stride {wildcards.stride}\
--d-embed {wildcards.embed}\
--n-heads {wildcards.heads}\
--budget-tol {config[budget_tolerance]}\
--store {input.store}\
--out-dir {output}\
--run-id {wildcards.arch}_{wildcards.params}_s{wildcards.stride}_g{wildcards.classes}\
--stage ladder_point\
--seed {config[seed]}
"""
rule train_control:
input:
data_store='../data/ml_inputs/{model}_data_store'
output:
directory('../data/ml_models/control_run__{model}')
conda:
config["torch_env"]
shell:
"""
python -m custom_models train\
--stage control
"""
+2 -2
View File
@@ -60,13 +60,13 @@ rule download_reference_genome:
rule concatenate_reference_genomes:
input:
expand("../data/reference_genomes/{species}.fasta",species=config["species_list"])
expand("../data/reference_genomes/{species}.fasta",species=config["species_list"]+config["trap_species"])
output:
"../data/reference_genomes/full_reference.fasta"
threads: 1
shell:
"""
cat {input} > {output}
cat {input} {config[ont_dcs_fasta]} > {output}
"""
rule create_minimap_index: