Files
eDNA_Stream_E01/workflow/rules/ml_pipeline.smk
T

66 lines
1.9 KiB
Plaintext

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
"""