import logging
from pathlib import Path
from typing import Dict, List, Union
from ewokstools.submit import save_and_execute, wait_to_finish_queue
from ewoksutils.task_utils import task_inputs
from ._utils import (
SLURM_JOB_PARAMETERS_INFERENCE,
EwoksWorkflow,
_confirm_submission,
_get_destination_filename,
_normalize_slurm_job_parameters,
_print_checking_template_format,
_print_job_progress,
_print_number_of_jobs_to_submit,
_print_template_invalid,
_print_template_valid,
_transfer_cif_files,
_validate_yaml_template_model,
_warning_dry_run_mode,
_warning_if_too_many_jobs,
get_list_scan_info_from_bliss_filenames_id13,
)
from .resources import INFERENCE_WORKFLOW
from .resources.models import InferenceModel
logger = logging.getLogger(__name__)
# def get_inputs_training(**kwargs):
# # TODO give interface to training inference NN
# nxdata_url = kwargs.get("nxdata_url")
# if nxdata_url:
# nxprocess_path_integrate = nxdata_url.split("::")[-1]
# else:
# file_input = kwargs.get("processed_data_filenames", [])[0]
# nxprocess_path_integrate = kwargs.get("nxprocess_path_integrate")
# nxdata_url = f"{file_input}::{nxprocess_path_integrate}/integrated"
# inputs_dict_training = {
# "nxdata_url": nxdata_url,
# "references_directory": kwargs.get("references_directory"),
# "wavelength": kwargs.get("wavelength"),
# "radial_limits": kwargs.get("radial_limits"),
# "submit_to_slurm": False,
# "inference_weights_filename": kwargs.get("inference_weights_filename"),
# }
# return validate_inputs_ewoks(
# inputs=inputs_dict_training,
# ewoks_task=PhaseInferenceTrainModel,
# id="training",
# )
[docs]
def main_inference(args):
for file in args.FILES:
if file.endswith((".yaml", ".yml")):
_main_inference_from_template(filename=file)
else:
logger.warning(
f"File {file} has an unsupported extension. Skipping it. Supported extensions are .yaml and .yml."
)
def _main_inference_from_template(filename: Union[str, Path]):
_print_checking_template_format()
is_valid, result = _validate_yaml_template_model(filename, InferenceModel)
if not is_valid:
_print_template_invalid(filename, result)
return
_print_template_valid()
_main_inference(**result)
def _main_inference(
processed_filenames: Union[str, List[str]],
nxprocess_path_filtered: str,
nxprocess_path_inference: str,
submit_parameters: Dict = None,
**kwargs,
):
workflow_inference = EwoksWorkflow(INFERENCE_WORKFLOW)
submit_parameters = submit_parameters or {}
dry_run = not submit_parameters.get("submit")
list_scans_info = get_list_scan_info_from_bliss_filenames_id13(
bliss_filenames=processed_filenames,
**kwargs,
)
nb_total_jobs = len(list_scans_info)
_print_number_of_jobs_to_submit(nb_total_jobs)
_confirm_submission()
if dry_run:
_warning_dry_run_mode()
else:
_warning_if_too_many_jobs(nb_total_jobs)
submitted_jobs = []
for index_job, scan_info in enumerate(list_scans_info):
output_filename = scan_info.filename_raw_dataset
external_output_filename = output_filename
inference_parameters = kwargs.get("inference", {}).copy()
wavelength_A = inference_parameters.pop("wavelength")
radial_limits = inference_parameters.pop("radial_limits")
references_directory = inference_parameters.pop("references_directory")
nxprocess_name_inference = nxprocess_path_inference.split("/")[-1]
target_cif_directory = _transfer_cif_files(
references_directory=str(Path(references_directory).resolve()),
target_directory=Path(external_output_filename).parent
/ f".{nxprocess_name_inference}_cif",
)
inputs_ewoks = task_inputs(
id=workflow_inference.inference.id,
task_identifier=workflow_inference.inference.task_identifier,
inputs={
"nxdata_url": f"{output_filename}::{nxprocess_path_filtered}/filtered",
"wavelength": wavelength_A,
"nxprocess_name": nxprocess_name_inference,
"radial_limits": radial_limits,
"destination_file": _get_destination_filename(
output_filename, scan_info.scan_nb, nxprocess_name_inference
),
"worker_module": "scattering",
"output_filename": output_filename,
"external_output_filename": external_output_filename,
"submit_to_slurm": False,
"references_directory": target_cif_directory,
**inference_parameters,
},
)
submit_parameters_ = submit_parameters.copy()
submitted = save_and_execute(
workflow=INFERENCE_WORKFLOW,
inputs=inputs_ewoks,
destination_filename=_get_destination_filename(
output_filename, scan_info.scan_nb, nxprocess_name_inference
),
slurm_job_parameters=_normalize_slurm_job_parameters(
SLURM_JOB_PARAMETERS_INFERENCE,
submit_parameters_.pop("slurm_job_parameters", {}),
),
**submit_parameters_,
)
submitted_jobs.append(submitted)
_print_job_progress(index_job + 1, nb_total_jobs, dry_run)
wait_to_finish_queue(submitted_jobs)