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
+3
View File
@@ -0,0 +1,3 @@
[submodule "custom_models"]
path = custom_models
url = git@git.tk-ai.eu:Tom/eDNA_Stream_E01_custom_models.git
+8 -4
View File
@@ -15,12 +15,16 @@
- python 3.14
- snakemake
- python-dotenv
## Optional:
- custom-models package: live in a separate repo (`git@git.tk-ai.eu:Tom/eDNA_Stream_E01_custom_models.git`) included here as the `custom_models` git submodule; the workflow pip-installs it from that checkout (see `workflow/envs/torch_*.yaml`)
- clone with: `git clone --recurse-submodules ...` (or `git submodule update --init` on an existing clone)
- update to a newer version: `git submodule update --remote custom_models`, then commit the new pinned commit
- .env file with NCBI API key -> slightly speeds up reference fasta download
## Running:
```bash
cd workflow
snakemake --cores <cores> --resources gpu=<n_gpus>
```
snakemake --cores <cores> --resources gpu=<n_gpus> --use-conda
```
## Notes:
- Reads aligned to the FP traps by minimap are excluded for training -> esp. the E. coli reads are highly represented and should not be there, but are stable even when only using stringent alignment
Submodule
+1
Submodule custom_models added at 96ae5d78ba
+2
View File
@@ -0,0 +1,2 @@
>DCS_Lambda DCS Lambda
GCCATCAGATTGTGTTTGTTAGTCGCTTTTTTTTTTTGGAATTTTTTTTTTGGAATTTTTTTTTTGCGCTAACAACCTCCTGCCGTTTTGCCCGTGCATATCGGTCACGAACAAATCTGATTACTAAACACAGTAGCCTGGATTTGTTCTATCAGTAATCGACCTTATTCCTAATTAAATAGAGCAAATCCCCTTATTGGGGGTAAGACATGAAGATGCCAGAAAAACATGACCTGTTGGCCGCCATTCTCGCGGCAAAGGAACAAGGCATCGGGGCAATCCTTGCGTTTGCAATGGCGTACCTTCGCGGCAGATATAATGGCGGTGCGTTTACAAAAACAGTAATCGACGCAACGATGTGCGCCATTATCGCCTAGTTCATTCGTGACCTTCTCGACTTCGCCGGACTAAGTAGCAATCTCGCTTATATAACGAGCGTGTTTATCGGCTACATCGGTACTGACTCGATTGGTTCGCTTATCAAACGCTTCGCTGCTAAAAAAGCCGGAGTAGAAGATGGTAGAAATCAATAATCAACGTAAGGCGTTCCTCGATATGCTGGCGTGGTCGGAGGGAACTGATAACGGACGTCAGAAAACCAGAAATCATGGTTATGACGTCATTGTAGGCGGAGAGCTATTTACTGATTACTCCGATCACCCTCGCAAACTTGTCACGCTAAACCCAAAACTCAAATCAACAGGCGCCGGACGCTACCAGCTTCTTTCCCGTTGGTGGGATGCCTACCGCAAGCAGCTTGGCCTGAAAGACTTCTCTCCGAAAAGTCAGGACGCTGTGGCATTGCAGCAGATTAAGGAGCGTGGCGCTTTACCTATGATTGATCGTGGTGATATCCGTCAGGCAATCGACCGTTGCAGCAATATCTGGGCTTCACTGCCGGGCGCTGGTTATGGTCAGTTCGAGCATAAGGCTGACAGCCTGATTGCAAAATTCAAAGAAGCGGGCGGAACGGTCAGAGAGATTGATGTATGAGCAGAGTCACCGCGATTATCTCCGCTCTGGTTATCTGCATCATCGTCTGCCTGTCATGGGCTGTTAATCATTACCGTGATAACGCCATTACCTACAAAGCCCAGCGCGACAAAAATGCCAGAGAACTGAAGCTGGCGAACGCGGCAATTACTGACATGCAGATGCGTCAGCGTGATGTTGCTGCGCTCGATGCAAAATACACGAAGGAGTTAGCTGATGCTAAAGCTGAAAATGATGCTCTGCGTGATGATGTTGCCGCTGGTCGTCGTCGGTTGCACATCAAAGCAGTCTGTCAGTCAGTGCGTGAAGCCACCACCGCCTCCGGCGTGGATAATGCAGCCTCCCCCCGACTGGCAGACACCGCTGAACGGGATTATTTCACCCTCAGAGAGAGGCTGATCACTATGCAAAAACAACTGGAAGGAACCCAGAAGTATATTAATGAGCAGTGCAGATAGAGTTGCCCATATCGATGGGCAACTCATGCAATTATTGTGAGCAATACACACGCGCTTCCAGCGGAGTATAAATGCCTAAAGTAATAAAACCGAGCAATCCATTTACGAATGTTTGCTGGGTTTCTGTTTTAACAACATTTTCTGCGCCGCCACAAATTTTGGCTGCATCGACAGTTTTCTTCTGCCCAATTCCAGAAACGAAGAAATGATGGGTGATGGTTTCCTTTGGTGCTACTGCTGCCGGTTTGTTTTGAACAGTAAACGTCTGTTGAGCACATCCTGTAATAAGCAGGGCCAGCGCAGTAGCGAGTAGCATTTTTTTCATGGTGTTATTCCCGATGCTTTTTGAAGTTCGCAGAATCGTATGTGTAGAAAATTAAACAAACCCTAAACAATGAGTTGAAATTTCATATTGTTAATATTTATTAATGTATGTCAGGTGCGATGAATCGTCATTGTATTCCCGGATTAACTATGTCCACAGCCCTGACGGGGAACTTCTCTGCGGGAGTGTCCGGGAATAATTAAAACGATGCACACAGGGTTTAGCGCGTACACGTATTGCATTATGCCAACGCCCCGGTGCTGACACGGAAGAAACCGGACGTTATGATTTAGCGTGGAAAGATTTGTGTAGTGTTCTGAATGCTCTCAGTAAATAGTAATGAATTATCAAAGGTATAGTAATATCTTTTATGTTCATGGATATTTGTAACCCATCGGAAAACTCCTGCTTTAGCAAGATTTTCCCTGTATTGCTGAAATGTGATTTCTCTTGATTTCAACCTATCATAGGACGTTTCTATAAGATGCGTGTTTCTTGAGAATTTAACATTTACAACCTTTTTAAGTCCTTTTATTAACACGGTGTTATCGTTTTCTAACACGATGTGAATATTATCTGTGGCTAGATAGTAAATATAATGTGAGACGTTGTGACGTTTTAGTTCAGAATAAAACAATTCACAGTCTAAATCTTTTCGCACTTGATCGAATATTTCTTTAAAAATGGCAACCTGAGCCATTGGTAAAACCTTCCATGTGATACGAGGGCGCGTAGTTTGCATTATCGTTTTTATCGTTTCAATCTGGTCTGACCTCCTTGTGTTTTGTTGATGATTTATGTCAAATATTAGGAATGTTTTCACTTAATAGTATTGGTTGCGTAACAAAGTGCGGTCCTGCTGGCATTCTGGAGGGAAATACAACCGACAGATGTATGTAAGGCCAACGTGCTCAAATCTTCATACAGAAAGATTTGAAGTAATATTTTAACCGCTAGATGAAGAGCAAGCGCATGGAGCGACAAAATGAATAAAGAACAATCTGCTGATGATCCCTCCGTGGATCTGATTCGTGTAAAAAATATGCTTAATAGCACCATTTCTATGAGTTACCCTGATGTTGTAATTGCATGTATAGAACATAAGGTGTCTCTGGAAGCATTCAGAGCAATTGAGGCAGCGTTGGTGAAGCACGATAATAATATGAAGGATTATTCCCTGGTGGTTGACTGATCACCATAACTGCTAATCATTCAAACTATTTAGTCTGTGACAGAGCCAACACGCAGTCTGTCACTGTCAGGAAAGTGGTAAAACTGCAACTCAATTACTGCAATGCCCTCGTAATTAAGTGAATTTACAATATCGTCCTGTTCGGAGGGAAGAACGCGGGATGTTCATTCTTCATCACTTTTAATTGATGTATATGCTCTCTTTTCTGACGTTAGTCTCCGACGGCAGGCTTCAATGACCCAGGCTGAGAAATTCCCGGACCCTTTTTGCTCAAGAGCGATGTTAATTTGTTCAATCATTTGGTTAGGAAAGCGGATGTTGCGGGTTGTTGTTCTGCGGGTTCTGTTCTTCGTTGACATGAGGTTGCCCCGTATTCAGTGTCGCTGATTTGTATTGTCTGAAGTTGTTTTTACGTTAAGTTGATGCAGATCAATTAATACGATACCTGCGTCATAATTGATTATTTGACGTGGTTTGATGGCCTCCACGCACGTTGTGATATGTAGATGATAATCATTATCACTTTACGGGTCCTTTCCGGTGAAAAAAAAGGTACCAAAAAAAACATCGTCGTGAGTAGTGAACCGTAAGC
+5 -2
View File
@@ -3,8 +3,6 @@ species_list:
- bacillus_subtilis
- listeria_monocytogenes
- staphylococcus_aureus
# Phylum Bacillota / Firmicutes (False-Positive Checkpoint - Near Bacillus)
- paenibacillus_polymyxa
# Phylum Pseudomonadota / Proteobacteria (Represented)
- cronobacter_sakazakii
- citrobacter_freundii
@@ -14,6 +12,11 @@ species_list:
- salmonella_enterica
- shigella_flexneri
- vibrio_cholerae
trap_species:
# Phylum Pseudomonadota / Proteobacteria (False-Positive Checkpoints - Near Enterobacteriaceae & Vibrio)
- escherichia_coli_k-12
unused_trap_species:
# Phylum Pseudomonadota / Proteobacteria (False-Positive Checkpoints - Near Enterobacteriaceae & Vibrio)
- aeromonas_hydrophila
# Phylum Bacillota / Firmicutes (False-Positive Checkpoint - Near Bacillus)
- paenibacillus_polymyxa
+8 -2
View File
@@ -15,6 +15,9 @@ module baseline:
snakefile: "rules/baseline_pipeline.smk"
config: config
module ml:
snakefile: 'rules/ml_pipeline.smk'
config: config
# Pulldown
rule all:
input:
@@ -26,8 +29,11 @@ rule all:
#'../data/pod5_files_to_pull',
#'../data/basecalled_reads/hac.fastq.gz'
#expand('../data/raw_pod5/PBK98658_853a956f_57f83f46_{batch}.pod5',batch=range(1,config["pod5_dataset_size"]+1,config["pod5_stride"]))
'../data/aligned_reads/{model}_to_genome.sorted.bam'
'../data/aligned_reads/nomiss_hac_to_genome.sorted.bam',
'../data/aligned_reads/trap_hac_to_genome.sorted.bam'
#'../data/ml_inputs/hac_label_store.pq'
#'../data/ml_inputs/model_layouts/cnn_512_4_100000_11_4.json'
use rule * from preparation
use rule * from baseline
use rule * from ml
BIN
View File
Binary file not shown.
+1
View File
@@ -1,3 +1,4 @@
arch: "amd64"
datasets_binary : "binaries/datasets"
ont_dcs_fasta : '../reference/ont_control_sequence.fasta'
+8 -1
View File
@@ -1,5 +1,7 @@
tmp_dir: "/tmp"
seed : 0
# dorado
dorado_model_dir : "../dorado_models"
hac_model_name: "dna_r10.4.1_e8.2_400bps_hac@v6.0.0"
@@ -13,4 +15,9 @@ fast_min_q: 8 # empty string or 0 to disable
# squiqqle dataset
pod5_dataset: "nomiss_96BC_P2I_SUP_2026"
pod5_dataset_size : 293
pod5_stride : 12 # Use 1 in x pod5 files from the ONT NO-MISS dataset to reduce dataset size / compute requirements
pod5_stride : 600 # Use 1 in x pod5 files from the ONT NO-MISS dataset to reduce dataset size / compute requirements
# custom ML
torch_env : '../envs/torch_cpu.yaml'
budget_tolerance : 0.1
reads_per_genus : 5000
+10
View File
@@ -0,0 +1,10 @@
channels:
- conda-forge
- bioconda
dependencies:
- python=3.14
- pysam=0.24
- pandas=3.0
- pyarrow
- biopython
- pydantic
+1
View File
@@ -1,3 +1,4 @@
name: samtools-test
channels:
- bioconda
dependencies:
+17
View File
@@ -0,0 +1,17 @@
name: e1-torch-cpu
channels:
- conda-forge
dependencies:
- python=3.14
- pytorch=2.13.*=cpu*
- numpy>=2.0,<3
- polars>=1.0
- pip
- pytest>=8
- pip:
- pod5
# E1 genus probe package, checked out as the `custom_models` submodule
# in the repo root (torch comes from conda here and satisfies the
# requirement, so only this package gets installed).
# Relative path assumes snakemake is invoked from `workflow/`.
- ../custom_models
+15
View File
@@ -0,0 +1,15 @@
name: e1-torch-cuda
channels:
- conda-forge
dependencies:
- python=3.14
- pytorch=2.13.*=cuda129*
- numpy>=2.0,<3
- polars>=1.0
- pip
- pip:
- pod5
# E1 genus probe package, checked out as the `custom_models` submodule
# in the repo root (torch comes from conda here).
# Relative path assumes snakemake is invoked from `workflow/`.
- ../custom_models
+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:
+91
View File
@@ -0,0 +1,91 @@
import pysam
from Bio import SeqIO
import pandas as pd
from pydantic import BaseModel
class ReferenceMetadata(BaseModel):
ref_id : str
genus : str
species : str
ref_len : int
if __name__ == '__main__':
input_bamfile = snakemake.input[0]
reference_genomes = snakemake.input[1]
output_parquet = snakemake.output[0]
# Get Reference Sequence ID -> Taxonomy mapping
ref_tax_data = {}
for ref in SeqIO.parse(reference_genomes,'fasta'):
ref_tax_data[ref.id] = {
'genus' : ref.description.split(' ')[1].lower(),
'species' : '_'.join(ref.description.split(' ')[1:3]).lower(),
'ref_length' : len(ref.seq)
}
# Open Bamfile for reading
bamfile = pysam.AlignmentFile(input_bamfile,'rb')
# Transcribe Taxonomy from Sequnce ID dict to bam header position list for fast lookup
ref_data = []
for ref_id in bamfile.header.references:
ref_data.append(
ReferenceMetadata(
ref_id = ref_id,
genus = ref_tax_data[ref_id]['genus'],
species = ref_tax_data[ref_id]['species'],
ref_len = ref_tax_data[ref_id]['ref_length']
)
)
# Parse alignments
data = []
n_records = 0
n_unmapped = 0
n_secondary = 0
n_supplementary = 0
n_primary = 0
for record in bamfile.fetch():
n_records += 1
if record.is_unmapped:
n_unmapped += 1
continue
if record.is_secondary:
n_secondary += 1
continue
if record.is_supplementary:
n_supplementary += 1
continue
n_primary += 1
nm_count = -1
for tag in record.get_tags():
if tag[0] == 'NM':
nm_count = tag[1]
assert nm_count != -1
assert isinstance(nm_count,int)
rlen = record.infer_read_length()
assert rlen is not None
qlen = record.infer_query_length()
assert qlen is not None
ref = ref_data[record.reference_id]
data.append(
{
'read_id' : record.query_name,
'ref_name' : ref.ref_id,
'genus' : ref.genus,
'species' : ref.species,
'identity' : 1 - (nm_count / qlen),
'aligned_len' : qlen,
'query_len' : rlen,
}
)
df = pd.DataFrame(data)
df.to_parquet(output_parquet,index=False)
print(f'Processed {input_bamfile}:')
print(f'\t{n_records} alignment records')
print(f'\t{n_primary} primary alignments')
print(f'\t{n_unmapped} unmapped reads')
print(f'\t{n_secondary+n_supplementary} non-primary alignments')
print(f'\t{df.shape[0]} alignment records written to disk')
print('Per-Genus counts:')
print(df.groupby('genus').aggregate(n_reads=('read_id','nunique')).reset_index())