Recursively change the interpolation mode in the applied operation stacks, default to "nearest". See also: :py:class:`monai.transform.inverse.InvertibleTransform` Args: trans_info: applied operation stack, tracking the previously applied invertible transform. mode: tar
(trans_info, mode: str = "nearest", align_corners: bool | None = None)
| 1767 | |
| 1768 | |
| 1769 | def convert_applied_interp_mode(trans_info, mode: str = "nearest", align_corners: bool | None = None): |
| 1770 | """ |
| 1771 | Recursively change the interpolation mode in the applied operation stacks, default to "nearest". |
| 1772 | |
| 1773 | See also: :py:class:`monai.transform.inverse.InvertibleTransform` |
| 1774 | |
| 1775 | Args: |
| 1776 | trans_info: applied operation stack, tracking the previously applied invertible transform. |
| 1777 | mode: target interpolation mode to convert, default to "nearest" as it's usually used to save the mode output. |
| 1778 | align_corners: target align corner value in PyTorch interpolation API, need to align with the `mode`. |
| 1779 | |
| 1780 | """ |
| 1781 | if isinstance(trans_info, (list, tuple)): |
| 1782 | return [convert_applied_interp_mode(x, mode=mode, align_corners=align_corners) for x in trans_info] |
| 1783 | if not isinstance(trans_info, Mapping): |
| 1784 | return trans_info |
| 1785 | trans_info = dict(trans_info) |
| 1786 | if "mode" in trans_info: |
| 1787 | current_mode = trans_info["mode"] |
| 1788 | if isinstance(current_mode, int) or current_mode in _interp_modes: |
| 1789 | trans_info["mode"] = mode |
| 1790 | elif isinstance(current_mode[0], int) or current_mode[0] in _interp_modes: |
| 1791 | trans_info["mode"] = [mode for _ in range(len(mode))] |
| 1792 | if "align_corners" in trans_info: |
| 1793 | _align_corners = TraceKeys.NONE if align_corners is None else align_corners |
| 1794 | current_value = trans_info["align_corners"] |
| 1795 | trans_info["align_corners"] = ( |
| 1796 | [_align_corners for _ in mode] if issequenceiterable(current_value) else _align_corners |
| 1797 | ) |
| 1798 | if ("mode" not in trans_info) and ("align_corners" not in trans_info): |
| 1799 | return { |
| 1800 | k: convert_applied_interp_mode(trans_info[k], mode=mode, align_corners=align_corners) for k in trans_info |
| 1801 | } |
| 1802 | return trans_info |
| 1803 | |
| 1804 | |
| 1805 | def reset_ops_id(data): |
searching dependent graphs…