Skip to content

nabu.io.crop_volume

[docs] module nabu.io.crop_volume

  1
  2
  3
  4
  5
  6
  7
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
from __future__ import annotations


import os
from typing import TYPE_CHECKING, Literal

from silx.io.url import DataUrl

from tomoscan.volumebase import VolumeBase
from tomoscan.esrf.volume.singleframebase import VolumeSingleFrameBase
from tomoscan.esrf.volume.hdf5volume import HDF5Volume
from tomoscan.esrf.volume.tiffvolume import MultiTIFFVolume
from ..io.writer import get_datetime
from .. import version as nabu_version

if TYPE_CHECKING:
    from collections.abc import Callable

    from tqdm import tqdm

    from .models.crop_volume import CropModel


def _get_output_volume_shape(model: CropModel, input_volume: VolumeBase) -> tuple[int, int, int]:
    """compute the output volume shape according to the cropping model and the input volume shape"""

    shape = input_volume.get_volume_shape()

    def get_bound(axis: Literal[0, 1, 2]):
        if axis == 0:
            current_slice = model.axis_0_slice
        elif axis == 1:
            current_slice = model.axis_1_slice
        elif axis == 2:
            current_slice = model.axis_2_slice
        else:
            raise ValueError(f"axis should be 0, 1 or 2. Got {axis}")
        if current_slice is None:
            return shape[axis]
        start = current_slice.start
        stop = current_slice.stop if current_slice.stop is not None else shape[axis]
        return stop - start

    return tuple([get_bound(i) for i in range(3)])


def crop_volume(input_volume: VolumeBase, output_volume: VolumeBase, model: CropModel, progress: tqdm | None = None):
    for volume, volume_type in {input_volume: "input_volume", output_volume: "output_volume"}.items():
        if not isinstance(volume, VolumeBase):
            raise TypeError(f"'{volume_type}' should be of type VolumeBase. Got {type(volume)}")

    frame_dumper_generator = output_volume.data_file_saver_generator(
        n_frames=_get_output_volume_shape(model=model, input_volume=input_volume)[0],
        data_url=output_volume.data_url,
        overwrite=output_volume.overwrite,
    )
    if progress is not None:
        progress.total = input_volume.get_volume_shape()[0]
    for i_slice, input_slice in enumerate(
        input_volume.browse_slices(),
    ):
        if progress is not None:
            progress.update()
        # note: browsing is done along the axis 0
        if model.axis_0_slice is not None:
            frame_is_before_start = i_slice < model.axis_0_slice.start
            frame_is_after_start = model.axis_0_slice.stop != -1 and i_slice >= model.axis_0_slice.stop
            if frame_is_before_start or frame_is_after_start:
                # frame is ignored if the defined ROI along axis 0 does not include it
                continue

        output_slice = input_slice
        if model.axis_1_slice is not None:
            output_slice = output_slice[model.axis_1_slice]
        if model.axis_2_slice is not None:
            output_slice = output_slice[:, model.axis_2_slice]

        frame_dumper = next(frame_dumper_generator)
        frame_dumper[:] = output_slice

    model_dict = model.model_dump()
    volume.metadata = {
        "crop_volume": {
            "configuration": model_dict,
            "input_volume": input_volume.get_identifier().to_str(),
            "program": "nabu-crop",
            "date": get_datetime(),
            "version": nabu_version,
        }
    }
    volume.save_metadata()


def get_default_output_volume(input_volume: VolumeBase) -> VolumeBase:
    if not isinstance(input_volume, VolumeBase):
        raise TypeError(f"input_volume is expected to be an instance of {VolumeBase}")

    def _update_folder_path(dir_name):
        parent_dir_name, folder = os.path.split(dir_name)
        return os.path.join(
            parent_dir_name,
            folder + "_cropped",
        )

    def _update_file_path(dir_name):
        parent_dir_name, file = os.path.split(dir_name)
        file_basename, file_ext = os.path.splitext(file)
        return os.path.join(
            parent_dir_name,
            file_basename + "_cropped" + file_ext,
        )

    def _update_url(url, update_fct: Callable[[str], str]):
        return DataUrl(
            file_path=update_fct(url.file_path()),
            data_path=url.data_path(),
            scheme=url.scheme(),
            data_slice=url.data_slice(),
        )

    if isinstance(input_volume, VolumeSingleFrameBase):

        new_data_url = _update_url(input_volume.data_url, update_fct=_update_folder_path)
        new_metadata_url = _update_url(input_volume.metadata_url, update_fct=_update_folder_path)
        return type(input_volume)(
            data_url=new_data_url,
            metadata_url=new_metadata_url,
        )

    elif isinstance(input_volume, (HDF5Volume, MultiTIFFVolume)):
        new_data_url = _update_url(input_volume.data_url, update_fct=_update_file_path)
        new_metadata_url = _update_url(input_volume.metadata_url, update_fct=_update_file_path)
        return type(input_volume)(
            data_url=new_data_url,
            metadata_url=new_metadata_url,
        )
    else:
        raise NotImplementedError(f"input volume format {type(input_volume)} is not handled")