Source code for ewoksid13.scripts.inference

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)