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