diff --git a/workflow/Snakefile b/workflow/Snakefile index 50625da..01747bd 100644 --- a/workflow/Snakefile +++ b/workflow/Snakefile @@ -30,7 +30,7 @@ rule all: #'../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/nomiss_hac_to_genome.sorted.bam', - '../data/aligned_reads/trap_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' diff --git a/workflow/config/main.yaml b/workflow/config/main.yaml index 1a2cf10..b040d7d 100644 --- a/workflow/config/main.yaml +++ b/workflow/config/main.yaml @@ -2,6 +2,9 @@ tmp_dir: "/tmp" seed : 0 +#gpu +gpu_stack : 'cuda' #'cuda' or 'rocm' + # dorado dorado_model_dir : "../dorado_models" hac_model_name: "dna_r10.4.1_e8.2_400bps_hac@v6.0.0" diff --git a/workflow/rules/baseline_pipeline.smk b/workflow/rules/baseline_pipeline.smk index 12ab7f2..983b64c 100644 --- a/workflow/rules/baseline_pipeline.smk +++ b/workflow/rules/baseline_pipeline.smk @@ -38,6 +38,8 @@ wildcard_constraints: batch=r"(batch)?\d+", model=r'hac|fast' +GPU_STACK = config.get('gpu_stack','cuda') + rule download_nomiss_pod5: output: '../data/raw_pod5/nomiss/PBK98658_853a956f_57f83f46_{batch}.pod5' @@ -53,7 +55,7 @@ rule download_trap_tar: temp(directory('../data/raw_pod5/trap/all')) threads: 32 params: - tmp='/tmp/trap_all.tar' + tmp='../data/raw_pod5/trap/trap_all.tar' shell: """ curl -L --http1.1 -A "Mozilla/5.0" "https://api.figshare.com/v2/file/download/45408628" -o {params.tmp} @@ -72,24 +74,41 @@ rule extract_trap_pod5: """ mv {params.input_pod5} {output} """ -rule basecall_pod5: - input: - pod5='../data/raw_pod5/{dataset}/{run}_{batch}.pod5', - model=get_model_requirement - output: - temp('../data/basecalled_reads/{dataset}/{model}/{run}_{batch}.fastq') - threads: - 32 - resources: - gpu=1 - params: - benchmarking=get_benchmarking_file, - min_qscore=get_qscore, - dorado_model=get_model_name, - shell: - """ - dorado basecaller --models-directory {config[dorado_model_dir]} --emit-fastq {params.benchmarking} {params.min_qscore} {params.dorado_model} {input.pod5} > {output} - """ + +if GPU_STACK == 'cuda': + rule basecall_pod5_dorado: + input: + pod5='../data/raw_pod5/{dataset}/{run}_{batch}.pod5', + model=get_model_requirement + output: + temp('../data/basecalled_reads/{dataset}/{model}/{run}_{batch}.fastq') + threads: + 32 + resources: + gpu=1 + params: + benchmarking=get_benchmarking_file, + min_qscore=get_qscore, + dorado_model=get_model_name, + shell: + """ + dorado basecaller --models-directory {config[dorado_model_dir]} --emit-fastq {params.benchmarking} {params.min_qscore} {params.dorado_model} {input.pod5} > {output} + """ +elif GPU_STACK == 'rocm': + rule basecall_pod5_slorado: + input: + pod5='../data/raw_pod5/{dataset}/{run}_{batch}.pod5', + model=get_model_requirement + output: + temp('../data/basecalled_reads/{dataset}/{model}/{run}_{batch}.fastq') + threads: + 32 + resources: + gpu=1 + shell: + """ + pod5-slorado basecaller -o {output} -t {threads} --flash=yes {input.model} {input.pod5} + """ rule concatenate_basecalled_fastq: input: