{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import argparse\nimport os\nimport sys\nimport numpy as np\nimport pandas as pd\nimport SimpleITK as sitk\nfrom scipy.ndimage import binary_dilation, binary_erosion\nfrom tqdm import tqdm\nimport csv\n\nimport torch\nfrom torch.cuda.amp import autocast\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom scipy.ndimage.interpolation import zoom\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-23T09:36:28.527431Z","iopub.execute_input":"2022-09-23T09:36:28.527923Z","iopub.status.idle":"2022-09-23T09:36:28.537322Z","shell.execute_reply.started":"2022-09-23T09:36:28.527865Z","shell.execute_reply":"2022-09-23T09:36:28.535728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\n\nclass Dict(dict):\n\n    def __init__(__self, *args, **kwargs):\n        object.__setattr__(__self, '__parent', kwargs.pop('__parent', None))\n        object.__setattr__(__self, '__key', kwargs.pop('__key', None))\n        object.__setattr__(__self, '__frozen', False)\n        for arg in args:\n            if not arg:\n                continue\n            elif isinstance(arg, dict):\n                for key, val in arg.items():\n                    __self[key] = __self._hook(val)\n            elif isinstance(arg, tuple) and (not isinstance(arg[0], tuple)):\n                __self[arg[0]] = __self._hook(arg[1])\n            else:\n                for key, val in iter(arg):\n                    __self[key] = __self._hook(val)\n\n        for key, val in kwargs.items():\n            __self[key] = __self._hook(val)\n\n    def __setattr__(self, name, value):\n        if hasattr(self.__class__, name):\n            raise AttributeError(\"'Dict' object attribute \"\n                                 \"'{0}' is read-only\".format(name))\n        else:\n            self[name] = value\n\n    def __setitem__(self, name, value):\n        isFrozen = (hasattr(self, '__frozen') and\n                    object.__getattribute__(self, '__frozen'))\n        if isFrozen and name not in super(Dict, self).keys():\n                raise KeyError(name)\n        super(Dict, self).__setitem__(name, value)\n        try:\n            p = object.__getattribute__(self, '__parent')\n            key = object.__getattribute__(self, '__key')\n        except AttributeError:\n            p = None\n            key = None\n        if p is not None:\n            p[key] = self\n            object.__delattr__(self, '__parent')\n            object.__delattr__(self, '__key')\n\n    def __add__(self, other):\n        if not self.keys():\n            return other\n        else:\n            self_type = type(self).__name__\n            other_type = type(other).__name__\n            msg = \"unsupported operand type(s) for +: '{}' and '{}'\"\n            raise TypeError(msg.format(self_type, other_type))\n\n    @classmethod\n    def _hook(cls, item):\n        if isinstance(item, dict):\n            return cls(item)\n        elif isinstance(item, (list, tuple)):\n            return type(item)(cls._hook(elem) for elem in item)\n        return item\n\n    def __getattr__(self, item):\n        return self.__getitem__(item)\n\n    def __missing__(self, name):\n        if object.__getattribute__(self, '__frozen'):\n            raise KeyError(name)\n        return self.__class__(__parent=self, __key=name)\n\n    def __delattr__(self, name):\n        del self[name]\n\n    def to_dict(self):\n        base = {}\n        for key, value in self.items():\n            if isinstance(value, type(self)):\n                base[key] = value.to_dict()\n            elif isinstance(value, (list, tuple)):\n                base[key] = type(value)(\n                    item.to_dict() if isinstance(item, type(self)) else\n                    item for item in value)\n            else:\n                base[key] = value\n        return base\n\n    def copy(self):\n        return copy.copy(self)\n\n    def deepcopy(self):\n        return copy.deepcopy(self)\n\n    def __deepcopy__(self, memo):\n        other = self.__class__()\n        memo[id(self)] = other\n        for key, value in self.items():\n            other[copy.deepcopy(key, memo)] = copy.deepcopy(value, memo)\n        return other\n\n    def update(self, *args, **kwargs):\n        other = {}\n        if args:\n            if len(args) > 1:\n                raise TypeError()\n            other.update(args[0])\n        other.update(kwargs)\n        for k, v in other.items():\n            if ((k not in self) or\n                (not isinstance(self[k], dict)) or\n                (not isinstance(v, dict))):\n                self[k] = v\n            else:\n                self[k].update(v)\n\n    def __getnewargs__(self):\n        return tuple(self.items())\n\n    def __getstate__(self):\n        return self\n\n    def __setstate__(self, state):\n        self.update(state)\n\n    def __or__(self, other):\n        if not isinstance(other, (Dict, dict)):\n            return NotImplemented\n        new = Dict(self)\n        new.update(other)\n        return new\n\n    def __ror__(self, other):\n        if not isinstance(other, (Dict, dict)):\n            return NotImplemented\n        new = Dict(other)\n        new.update(self)\n        return new\n\n    def __ior__(self, other):\n        self.update(other)\n        return self\n\n    def setdefault(self, key, default=None):\n        if key in self:\n            return self[key]\n        else:\n            self[key] = default\n            return default\n\n    def freeze(self, shouldFreeze=True):\n        object.__setattr__(self, '__frozen', shouldFreeze)\n        for key, val in self.items():\n            if isinstance(val, Dict):\n                val.freeze(shouldFreeze)\n\n    def unfreeze(self):\n        self.freeze(False)\n\ndef check_file_exist(filename, msg_tmpl='file \"{}\" does not exist'):\n    if not osp.isfile(filename):\n        raise FileNotFoundError(msg_tmpl.format(filename))","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.540321Z","iopub.execute_input":"2022-09-23T09:36:28.541174Z","iopub.status.idle":"2022-09-23T09:36:28.577188Z","shell.execute_reply.started":"2022-09-23T09:36:28.541131Z","shell.execute_reply":"2022-09-23T09:36:28.575425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Copyright (c) Open-MMLab. All rights reserved.\nimport ast\nimport os.path as osp\nimport platform\nimport shutil\nimport sys\nimport tempfile\nimport copy\n\nfrom argparse import Action, ArgumentParser\nfrom collections import abc\nfrom importlib import import_module\nfrom yapf.yapflib.yapf_api import FormatCode\n\nif platform.system() == 'Windows':\n    import regex as re\nelse:\n    import re\n\n\nBASE_KEY = '_base_'\nDELETE_KEY = '_delete_'\nRESERVED_KEYS = ['filename', 'text', 'pretty_text']\nclass ConfigDict(Dict):\n\n    def __missing__(self, name):\n        raise KeyError(name)\n\n    def __getattr__(self, name):\n        try:\n            value = super(ConfigDict, self).__getattr__(name)\n        except KeyError:\n            ex = AttributeError(f\"'{self.__class__.__name__}' object has no \"\n                                f\"attribute '{name}'\")\n        except Exception as e:\n            ex = e\n        else:\n            return value\n        raise ex\n\nclass Config:\n    \"\"\"A facility for config and config files.\n\n    It supports common file formats as configs: python/json/yaml. The interface\n    is the same as a dict object and also allows access config values as\n    attributes.\n\n    Example:\n        >>> cfg = Config(dict(a=1, b=dict(b1=[0, 1])))\n        >>> cfg.a\n        1\n        >>> cfg.b\n        {'b1': [0, 1]}\n        >>> cfg.b.b1\n        [0, 1]\n        >>> cfg = Config.fromfile('tests/data/config/a.py')\n        >>> cfg.filename\n        \"/home/kchen/projects/mmcv/tests/data/config/a.py\"\n        >>> cfg.item4\n        'test'\n        >>> cfg\n        \"Config [path: /home/kchen/projects/mmcv/tests/data/config/a.py]: \"\n        \"{'item1': [1, 2], 'item2': {'a': 0}, 'item3': True, 'item4': 'test'}\"\n    \"\"\"\n\n    @staticmethod\n    def _validate_py_syntax(filename):\n        with open(filename, 'r') as f:\n            content = f.read()\n        try:\n            ast.parse(content)\n        except SyntaxError as e:\n            raise SyntaxError('There are syntax errors in config '\n                              f'file {filename}: {e}')\n\n    @staticmethod\n    def _substitute_predefined_vars(filename, temp_config_name):\n        file_dirname = osp.dirname(filename)\n        file_basename = osp.basename(filename)\n        file_basename_no_extension = osp.splitext(file_basename)[0]\n        file_extname = osp.splitext(filename)[1]\n        support_templates = dict(\n            fileDirname=file_dirname,\n            fileBasename=file_basename,\n            fileBasenameNoExtension=file_basename_no_extension,\n            fileExtname=file_extname)\n        with open(filename, 'r') as f:\n            config_file = f.read()\n        for key, value in support_templates.items():\n            regexp = r'\\{\\{\\s*' + str(key) + r'\\s*\\}\\}'\n            value = value.replace('\\\\', '/')\n            config_file = re.sub(regexp, value, config_file)\n        with open(temp_config_name, 'w') as tmp_config_file:\n            tmp_config_file.write(config_file)\n\n    @staticmethod\n    def _file2dict(filename, use_predefined_variables=True):\n        filename = osp.abspath(osp.expanduser(filename))\n        check_file_exist(filename)\n        fileExtname = osp.splitext(filename)[1]\n        if fileExtname not in ['.py', '.json', '.yaml', '.yml']:\n            raise IOError('Only py/yml/yaml/json type are supported now!')\n\n        with tempfile.TemporaryDirectory() as temp_config_dir:\n            temp_config_file = tempfile.NamedTemporaryFile(\n                dir=temp_config_dir, suffix=fileExtname)\n            if platform.system() == 'Windows':\n                temp_config_file.close()\n            temp_config_name = osp.basename(temp_config_file.name)\n            # Substitute predefined variables\n            if use_predefined_variables:\n                Config._substitute_predefined_vars(filename,\n                                                   temp_config_file.name)\n            else:\n                shutil.copyfile(filename, temp_config_file.name)\n\n            if filename.endswith('.py'):\n                temp_module_name = osp.splitext(temp_config_name)[0]\n                sys.path.insert(0, temp_config_dir)\n                Config._validate_py_syntax(filename)\n                mod = import_module(temp_module_name)\n                sys.path.pop(0)\n                cfg_dict = {\n                    name: value\n                    for name, value in mod.__dict__.items()\n                    if not name.startswith('__')\n                }\n                # delete imported module\n                del sys.modules[temp_module_name]\n            elif filename.endswith(('.yml', '.yaml', '.json')):\n                import mmcv\n                cfg_dict = mmcv.load(temp_config_file.name)\n            # close temp file\n            temp_config_file.close()\n\n        cfg_text = filename + '\\n'\n        with open(filename, 'r') as f:\n            cfg_text += f.read()\n\n        if BASE_KEY in cfg_dict:\n            cfg_dir = osp.dirname(filename)\n            base_filename = cfg_dict.pop(BASE_KEY)\n            base_filename = base_filename if isinstance(\n                base_filename, list) else [base_filename]\n\n            cfg_dict_list = list()\n            cfg_text_list = list()\n            for f in base_filename:\n                _cfg_dict, _cfg_text = Config._file2dict(osp.join(cfg_dir, f))\n                cfg_dict_list.append(_cfg_dict)\n                cfg_text_list.append(_cfg_text)\n\n            base_cfg_dict = dict()\n            for c in cfg_dict_list:\n                if len(base_cfg_dict.keys() & c.keys()) > 0:\n                    raise KeyError('Duplicate key is not allowed among bases')\n                base_cfg_dict.update(c)\n\n            base_cfg_dict = Config._merge_a_into_b(cfg_dict, base_cfg_dict)\n            cfg_dict = base_cfg_dict\n\n            # merge cfg_text\n            cfg_text_list.append(cfg_text)\n            cfg_text = '\\n'.join(cfg_text_list)\n\n        return cfg_dict, cfg_text\n\n    @staticmethod\n    def _merge_a_into_b(a, b):\n        # merge dict `a` into dict `b` (non-inplace). values in `a` will\n        # overwrite `b`.\n        # copy first to avoid inplace modification\n        b = b.copy()\n        for k, v in a.items():\n            if isinstance(v, dict) and k in b and not v.pop(DELETE_KEY, False):\n                if not isinstance(b[k], dict):\n                    raise TypeError(\n                        f'{k}={v} in child config cannot inherit from base '\n                        f'because {k} is a dict in the child config but is of '\n                        f'type {type(b[k])} in base config. You may set '\n                        f'`{DELETE_KEY}=True` to ignore the base config')\n                b[k] = Config._merge_a_into_b(v, b[k])\n            else:\n                b[k] = v\n        return b\n\n    @staticmethod\n    def fromfile(filename, use_predefined_variables=True):\n        cfg_dict, cfg_text = Config._file2dict(filename,\n                                               use_predefined_variables)\n        return Config(cfg_dict, cfg_text=cfg_text, filename=filename)\n\n    @staticmethod\n    def auto_argparser(description=None):\n        \"\"\"Generate argparser from config file automatically (experimental)\"\"\"\n        partial_parser = ArgumentParser(description=description)\n        partial_parser.add_argument('config', help='config file path')\n        cfg_file = partial_parser.parse_known_args()[0].config\n        cfg = Config.fromfile(cfg_file)\n        parser = ArgumentParser(description=description)\n        parser.add_argument('config', help='config file path')\n        add_args(parser, cfg)\n        return parser, cfg\n\n    def __init__(self, cfg_dict=None, cfg_text=None, filename=None):\n        if cfg_dict is None:\n            cfg_dict = dict()\n        elif not isinstance(cfg_dict, dict):\n            raise TypeError('cfg_dict must be a dict, but '\n                            f'got {type(cfg_dict)}')\n        for key in cfg_dict:\n            if key in RESERVED_KEYS:\n                raise KeyError(f'{key} is reserved for config file')\n\n        super(Config, self).__setattr__('_cfg_dict', ConfigDict(cfg_dict))\n        super(Config, self).__setattr__('_filename', filename)\n        if cfg_text:\n            text = cfg_text\n        elif filename:\n            with open(filename, 'r') as f:\n                text = f.read()\n        else:\n            text = ''\n        super(Config, self).__setattr__('_text', text)\n\n    @property\n    def filename(self):\n        return self._filename\n\n    @property\n    def text(self):\n        return self._text\n\n    @property\n    def pretty_text(self):\n\n        indent = 4\n\n        def _indent(s_, num_spaces):\n            s = s_.split('\\n')\n            if len(s) == 1:\n                return s_\n            first = s.pop(0)\n            s = [(num_spaces * ' ') + line for line in s]\n            s = '\\n'.join(s)\n            s = first + '\\n' + s\n            return s\n\n        def _format_basic_types(k, v, use_mapping=False):\n            if isinstance(v, str):\n                v_str = f\"'{v}'\"\n            else:\n                v_str = str(v)\n\n            if use_mapping:\n                k_str = f\"'{k}'\" if isinstance(k, str) else str(k)\n                attr_str = f'{k_str}: {v_str}'\n            else:\n                attr_str = f'{str(k)}={v_str}'\n            attr_str = _indent(attr_str, indent)\n\n            return attr_str\n\n        def _format_list(k, v, use_mapping=False):\n            # check if all items in the list are dict\n            if all(isinstance(_, dict) for _ in v):\n                v_str = '[\\n'\n                v_str += '\\n'.join(\n                    f'dict({_indent(_format_dict(v_), indent)}),'\n                    for v_ in v).rstrip(',')\n                if use_mapping:\n                    k_str = f\"'{k}'\" if isinstance(k, str) else str(k)\n                    attr_str = f'{k_str}: {v_str}'\n                else:\n                    attr_str = f'{str(k)}={v_str}'\n                attr_str = _indent(attr_str, indent) + ']'\n            else:\n                attr_str = _format_basic_types(k, v, use_mapping)\n            return attr_str\n\n        def _contain_invalid_identifier(dict_str):\n            contain_invalid_identifier = False\n            for key_name in dict_str:\n                contain_invalid_identifier |= \\\n                    (not str(key_name).isidentifier())\n            return contain_invalid_identifier\n\n        def _format_dict(input_dict, outest_level=False):\n            r = ''\n            s = []\n\n            use_mapping = _contain_invalid_identifier(input_dict)\n            if use_mapping:\n                r += '{'\n            for idx, (k, v) in enumerate(input_dict.items()):\n                is_last = idx >= len(input_dict) - 1\n                end = '' if outest_level or is_last else ','\n                if isinstance(v, dict):\n                    v_str = '\\n' + _format_dict(v)\n                    if use_mapping:\n                        k_str = f\"'{k}'\" if isinstance(k, str) else str(k)\n                        attr_str = f'{k_str}: dict({v_str}'\n                    else:\n                        attr_str = f'{str(k)}=dict({v_str}'\n                    attr_str = _indent(attr_str, indent) + ')' + end\n                elif isinstance(v, list):\n                    attr_str = _format_list(k, v, use_mapping) + end\n                else:\n                    attr_str = _format_basic_types(k, v, use_mapping) + end\n\n                s.append(attr_str)\n            r += '\\n'.join(s)\n            if use_mapping:\n                r += '}'\n            return r\n\n        cfg_dict = self._cfg_dict.to_dict()\n        text = _format_dict(cfg_dict, outest_level=True)\n        # copied from setup.cfg\n        yapf_style = dict(\n            based_on_style='pep8',\n            blank_line_before_nested_class_or_def=True,\n            split_before_expression_after_opening_paren=True)\n        text, _ = FormatCode(text, style_config=yapf_style, verify=True)\n\n        return text\n\n    def __repr__(self):\n        return f'Config (path: {self.filename}): {self._cfg_dict.__repr__()}'\n\n    def __len__(self):\n        return len(self._cfg_dict)\n\n    def __getattr__(self, name):\n        return getattr(self._cfg_dict, name)\n\n    def __getitem__(self, name):\n        return self._cfg_dict.__getitem__(name)\n\n    def __setattr__(self, name, value):\n        if isinstance(value, dict):\n            value = ConfigDict(value)\n        self._cfg_dict.__setattr__(name, value)\n\n    def __setitem__(self, name, value):\n        if isinstance(value, dict):\n            value = ConfigDict(value)\n        self._cfg_dict.__setitem__(name, value)\n\n    def __iter__(self):\n        return iter(self._cfg_dict)\n\n    def __getstate__(self):\n        return (self._cfg_dict, self._filename, self._text)\n\n    def __setstate__(self, state):\n        _cfg_dict, _filename, _text = state\n        super(Config, self).__setattr__('_cfg_dict', _cfg_dict)\n        super(Config, self).__setattr__('_filename', _filename)\n        super(Config, self).__setattr__('_text', _text)\n\n    def dump(self, file=None):\n        cfg_dict = super(Config, self).__getattribute__('_cfg_dict').to_dict()\n        if self.filename.endswith('.py'):\n            if file is None:\n                return self.pretty_text\n            else:\n                with open(file, 'w') as f:\n                    f.write(self.pretty_text)\n        else:\n            import mmcv\n            if file is None:\n                file_format = self.filename.split('.')[-1]\n                return mmcv.dump(cfg_dict, file_format=file_format)\n            else:\n                mmcv.dump(cfg_dict, file)\n\n    def merge_from_dict(self, options):\n        \"\"\"Merge list into cfg_dict.\n\n        Merge the dict parsed by MultipleKVAction into this cfg.\n\n        Examples:\n            >>> options = {'model.backbone.depth': 50,\n            ...            'model.backbone.with_cp':True}\n            >>> cfg = Config(dict(model=dict(backbone=dict(type='ResNet'))))\n            >>> cfg.merge_from_dict(options)\n            >>> cfg_dict = super(Config, self).__getattribute__('_cfg_dict')\n            >>> assert cfg_dict == dict(\n            ...     model=dict(backbone=dict(depth=50, with_cp=True)))\n\n        Args:\n            options (dict): dict of configs to merge from.\n        \"\"\"\n        option_cfg_dict = {}\n        for full_key, v in options.items():\n            d = option_cfg_dict\n            key_list = full_key.split('.')\n            for subkey in key_list[:-1]:\n                d.setdefault(subkey, ConfigDict())\n                d = d[subkey]\n            subkey = key_list[-1]\n            d[subkey] = v\n\n        cfg_dict = super(Config, self).__getattribute__('_cfg_dict')\n        super(Config, self).__setattr__(\n            '_cfg_dict', Config._merge_a_into_b(option_cfg_dict, cfg_dict))","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.787139Z","iopub.execute_input":"2022-09-23T09:36:28.787554Z","iopub.status.idle":"2022-09-23T09:36:28.854898Z","shell.execute_reply.started":"2022-09-23T09:36:28.78752Z","shell.execute_reply":"2022-09-23T09:36:28.853417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dcm2nii(dcms_path):\n\t# 1.构建dicom序列文件阅读器，并执行（即将dicom序列文件“打包整合”）\n    reader = sitk.ImageSeriesReader()\n    dicom_names = reader.GetGDCMSeriesFileNames(dcms_path)\n    reader.SetFileNames(dicom_names)\n    image2 = reader.Execute()\n\t# 2.将整合后的数据转为array，并获取dicom文件基本信息\n    image_array = sitk.GetArrayFromImage(image2)  # z, y, x\n#     origin = image2.GetOrigin()  # x, y, z\n    spacing = image2.GetSpacing()  # x, y, z\n#     direction = image2.GetDirection()  # x, y, z\n# \t# 3.将array转为img，并保存为.nii.gz\n#     image3 = sitk.GetImageFromArray(image_array)\n#     image3.SetSpacing(spacing)\n#     image3.SetDirection(direction)\n#     image3.SetOrigin(origin)\n    return image_array, spacing","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.857712Z","iopub.execute_input":"2022-09-23T09:36:28.858125Z","iopub.status.idle":"2022-09-23T09:36:28.868902Z","shell.execute_reply.started":"2022-09-23T09:36:28.858082Z","shell.execute_reply":"2022-09-23T09:36:28.867397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegConfig:\n\n    def __init__(self,  network_f):\n        # TODO: 模型配置文件\n        self.network_f = network_f\n        if self.network_f is not None:\n\n            if isinstance(self.network_f, str):\n                self.network_cfg = Config.fromfile(self.network_f)\n            else:\n                import tempfile\n\n                with tempfile.TemporaryDirectory() as temp_config_dir:\n                    with tempfile.NamedTemporaryFile(dir=temp_config_dir, suffix='.py') as temp_config_file:\n                        with open(temp_config_file.name, 'wb') as f:\n                            f.write(self.network_f.read())\n\n                        self.network_cfg = Config.fromfile(temp_config_file.name)\n\n    def __repr__(self) -> str:\n        return str(self.__dict__)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.870701Z","iopub.execute_input":"2022-09-23T09:36:28.871105Z","iopub.status.idle":"2022-09-23T09:36:28.885858Z","shell.execute_reply.started":"2022-09-23T09:36:28.871072Z","shell.execute_reply":"2022-09-23T09:36:28.884338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegModel:\n\n    def __init__(self, model_f, network_f):\n        # TODO: 模型文件定制\n        self.model_f = model_f\n        self.network_f = network_f","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.888422Z","iopub.execute_input":"2022-09-23T09:36:28.889104Z","iopub.status.idle":"2022-09-23T09:36:28.897389Z","shell.execute_reply.started":"2022-09-23T09:36:28.889059Z","shell.execute_reply":"2022-09-23T09:36:28.895775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegPredictor:\n\n    def __init__(self, gpu: int, model: SegModel):\n        self.gpu = gpu\n        self.model = model\n        self.config = SegConfig(self.model.network_f)\n        self.load_model()\n\n    def load_model(self):\n        self.net = self._load_model(self.model.model_f, self.config.network_cfg, half=False)\n\n    def _load_model(self, model_f, network_f, half=False) -> None:\n        if isinstance(model_f, str):\n            # 根据后缀判断类型\n            if model_f.endswith(\".pth\"):\n                net = self.load_model_pth(model_f, network_f, half)\n            else:\n                net = self.load_model_jit(model_f, half)\n        else:\n            model_f.seek(0)\n            headers = model_f.peek(2)\n            if headers[0] == 0x80 and headers[1] == 0x02:\n                # pth文件类型\n                net = self.load_model_pth(model_f, network_f, half)\n            else:\n                # pt文件类型\n                net = self.load_model_jit(model_f, half)\n        return net\n\n    def load_model_pth(self, model_f, network_cfg, half) -> None:\n        # 加载动态图\n        config = network_cfg\n        print(config.model[\"backbone\"])\n        if config.model[\"backbone\"] == \"ResUnet\":\n            backbone = ResUnet(config.in_ch, channels=config.model[\"channels\"])\n        else:\n            raise TypeError(\"<<<<<<<<<<<wrong backbone>>>>>>>>>>>>\")\n\n        if config.model[\"head_type\"] == \"SegSoftHead\" or config.model[\"head_type\"] == \"SegSoftHead3D\":\n            head = SegSoftHead(config.model[\"channels\"], classes=8)\n        elif config.model[\"head_type\"]== \"SegSigHead\":\n            head = SegSigHead(config.model[\"channels\"], classes=1)\n        else:\n            raise TypeError(\"<<<<<<<<<<<wrong head>>>>>>>>>>>>\")\n            \n        net = SegNetwork(backbone, head)\n        checkpoint = torch.load(model_f, map_location=f\"cpu\")\n        if list(checkpoint[\"state_dict\"].keys())[0].startswith(\"module.\"):\n            state_dict = {k[7:]: v for k, v in checkpoint[\"state_dict\"].items()}\n        else:\n            state_dict = checkpoint[\"state_dict\"]\n        net.load_state_dict(state_dict, strict=False)\n        net.eval()\n        net.half()\n        net.cuda()\n        net = net.forward_test\n        return net\n\n    def forward(self, seg, spacing):\n\n        config = self.config.network_cfg\n        patch_size = np.array(config.patch_size)\n        ori_shape = np.array(seg.shape)\n        seg_np = zoom(seg, np.array(np.array(patch_size) / ori_shape), order=0)\n\n        data = torch.from_numpy(seg_np).float()[None, None]\n        # print(\"debug: \", np.unique(seg_np), torch.unique(data))\n        data = data.detach()\n        \n        with autocast():\n            data = data.cuda().detach()\n            pred_seg = self.net(data)\n            del data\n            if pred_seg.size()[1] > 1:\n                pred_seg = F.softmax(pred_seg, dim=1)\n                pred_seg = torch.argmax(pred_seg, dim=1, keepdim=True)\n                pred_seg[pred_seg == 8] = 0\n            else:\n                pred_seg = torch.sigmoid(pred_seg)\n                pred_seg[pred_seg >= config.threshold] = 1\n                pred_seg[pred_seg < config.threshold] = 0\n     \n        heatmap = pred_seg.cpu().detach().numpy()[0, 0].astype(np.uint8)\n        heatmap = zoom(heatmap, np.array(ori_shape / np.array(patch_size)), order=0)\n\n        heatmap *= (seg > 0)\n        heatmap = heatmap.astype(np.uint8)\n\n        return heatmap\n\nclass SegPredictor_:\n\n    def __init__(self, gpu: int, model: SegModel):\n        self.gpu = gpu\n        self.model = model\n        self.config = SegConfig(self.model.network_f)\n        self.load_model()\n\n    def load_model(self):\n        self.net = self._load_model(self.model.model_f, self.config.network_cfg, half=False)\n\n    def _load_model(self, model_f, network_f, half=False) -> None:\n        if isinstance(model_f, str):\n            net = self.load_model_pth(model_f, network_f, half)\n        else:\n            model_f.seek(0)\n            headers = model_f.peek(2)\n            net = self.load_model_pth(model_f, network_f, half)\n        return net\n    \n    def load_model_pth(self, model_f, network_cfg, half) -> None:\n        # 加载动态图\n        config = network_cfg\n        backbone = ResUnet(config.model[\"in_ch\"], channels=config.model[\"channels\"])\n\n        if config.model[\"head_type\"] == \"SegSoftHead\" or config.model[\"head_type\"] == \"SegSoftHead3D\":\n            head = SegSoftHead(in_channels=config.model[\"channels\"], classes=config.model[\"classes\"])\n        elif config.model[\"head_type\"] == \"SegSigHead\" or config.model[\"head_type\"] == \"SegSigHead3D\":\n            head = SegSigHead(in_channels=config.model[\"channels\"], classes=1)\n        else:\n            raise TypeError(\"<<<<<<<<<<<wrong head>>>>>>>>>>>>\")\n            \n        net = SegNetwork(backbone, head)\n        checkpoint = torch.load(model_f, map_location=f\"cpu\")\n        if list(checkpoint[\"state_dict\"].keys())[0].startswith(\"module.\"):\n            state_dict = {k[7:]: v for k, v in checkpoint[\"state_dict\"].items()}\n        else:\n            state_dict = checkpoint[\"state_dict\"]\n        net.load_state_dict(state_dict, strict=False)\n        # checkpoint = torch.load(model_f, map_location=f\"cpu\")\n        # net.load_state_dict(checkpoint[\"state_dict\"], strict=False)\n        net.eval()\n        net.half()\n        net.cuda()\n        net = net.forward_test\n        return net\n\n    def _get_input(self, vol, spacing_zyx):\n\n        config = self.config.network_cfg\n\n        def _window_array(vol, win_level, win_width):\n            win = [\n                win_level - win_width / 2,\n                win_level + win_width / 2,\n            ]\n            vol = torch.clamp(vol, win[0], win[1])\n            vol -= win[0]\n            vol /= win_width\n            return vol\n\n        vol = torch.from_numpy(vol).float()[None, None]\n        vol = [_window_array(vol, wl, wd) for wl, wd in zip(config.win_level, config.win_width)]\n        vol = torch.cat(vol, dim=1)\n        vol_shape = np.array(vol.shape[2:], dtype=np.float32)\n        patch_size = np.array(config.patch_size)\n        if np.any(vol_shape != patch_size):\n            vol = torch.nn.functional.interpolate(\n                vol.float(), size=tuple(patch_size), mode=\"trilinear\", align_corners=False\n            )\n        vol = vol.detach()\n        return vol\n\n    def forward(self, vol, spacing, sup_mask=None):\n\n        config = self.config.network_cfg\n        patch_size = np.array(config.patch_size)\n        ori_shape = np.array(vol.shape)\n        spacing_zyx = np.array(spacing)\n        data = self._get_input(vol, spacing_zyx)\n\n        # vol_save = data.squeeze().detach().cpu().numpy()\n        # vol_itk = sitk.GetImageFromArray(vol_save.astype(np.float32))\n        # sitk.WriteImage(vol_itk, f'./vol.nii.gz')\n\n        with autocast():\n            data = data.cuda().detach()\n            pred_seg = self.net(data)\n            del data\n            if pred_seg.size()[1] > 1:\n                pred_seg = F.softmax(pred_seg, dim=1)\n                pred_seg = torch.argmax(pred_seg, dim=1, keepdim=True)\n            else:\n                pred_seg = torch.sigmoid(pred_seg)\n                pred_seg[pred_seg >= config.threshold] = 1\n                pred_seg[pred_seg < config.threshold] = 0\n     \n        heatmap = pred_seg.cpu().detach().numpy()[0, 0].astype(np.uint8)\n        heatmap = zoom(heatmap, np.array(ori_shape / np.array(patch_size)), order=0)\n\n        heatmap = heatmap.astype(np.uint8)\n\n        return heatmap","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.902421Z","iopub.execute_input":"2022-09-23T09:36:28.903647Z","iopub.status.idle":"2022-09-23T09:36:28.94408Z","shell.execute_reply.started":"2022-09-23T09:36:28.903613Z","shell.execute_reply":"2022-09-23T09:36:28.94274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def conv3x3(in_planes, out_planes, stride=1, groups=1, dilation=1):\n    \"\"\"3x3 convolution with padding.\"\"\"\n    return nn.Conv3d(\n        in_planes,\n        out_planes,\n        kernel_size=3,\n        stride=stride,\n        padding=dilation,\n        groups=groups,\n        bias=False,\n        dilation=dilation,\n    )\n\ndef conv1x1(in_planes, out_planes, stride=1):\n    \"\"\"1x1 convolution.\"\"\"\n    return nn.Conv3d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.947717Z","iopub.execute_input":"2022-09-23T09:36:28.948325Z","iopub.status.idle":"2022-09-23T09:36:28.961261Z","shell.execute_reply.started":"2022-09-23T09:36:28.948276Z","shell.execute_reply":"2022-09-23T09:36:28.959821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None):\n        super(BasicBlock, self).__init__()\n\n        # Both self.conv1 and self.downsample layers downsample the input when stride != 1\n        self.conv1 = conv3x3(inplanes, planes, stride)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3(planes, planes)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        identity = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            identity = self.downsample(x)\n\n        out += identity\n\n        del identity\n        del x\n        torch.cuda.empty_cache()\n        \n        out = self.relu(out)\n\n        return out","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.963726Z","iopub.execute_input":"2022-09-23T09:36:28.964739Z","iopub.status.idle":"2022-09-23T09:36:28.977395Z","shell.execute_reply.started":"2022-09-23T09:36:28.964648Z","shell.execute_reply":"2022-09-23T09:36:28.975738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_res_layer(inplanes, planes, blocks, stride=1):\n    downsample = nn.Sequential(\n        conv1x1(inplanes, planes, stride),\n        nn.BatchNorm3d(planes),\n    )\n\n    layers = []\n    layers.append(BasicBlock(inplanes, planes, stride, downsample))\n    for _ in range(1, blocks):\n        layers.append(BasicBlock(planes, planes))\n\n    return nn.Sequential(*layers)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:28.980368Z","iopub.execute_input":"2022-09-23T09:36:28.981471Z","iopub.status.idle":"2022-09-23T09:36:28.995065Z","shell.execute_reply.started":"2022-09-23T09:36:28.981397Z","shell.execute_reply":"2022-09-23T09:36:28.993973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n\n    def __init__(self, in_ch, out_ch, stride=1, kernel_size=3):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv3d(in_ch, out_ch, kernel_size=kernel_size, stride=stride, padding=int(kernel_size / 2)),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_ch, out_ch, 3, padding=1, dilation=1),\n            nn.BatchNorm3d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, input):\n        return self.conv(input)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.033063Z","iopub.execute_input":"2022-09-23T09:36:29.0334Z","iopub.status.idle":"2022-09-23T09:36:29.044715Z","shell.execute_reply.started":"2022-09-23T09:36:29.03337Z","shell.execute_reply":"2022-09-23T09:36:29.04182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _ASPPConv(in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1):\n    asppconv = nn.Sequential(\n        nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=False),\n        nn.BatchNorm3d(out_channels),\n        nn.ReLU(inplace=True),\n    )\n    return asppconv\n\nclass ASPP(nn.Module):\n    \"\"\"\n    ASPP module in `DeepLabV3, see also in <https://arxiv.org/abs/1706.05587>` \n    \"\"\"\n    def __init__(self, in_channels, out_channels, output_stride=16):\n        super(ASPP, self).__init__()\n\n        if output_stride == 16:\n            astrous_rates = [0, 4, 8, 12]\n        elif output_stride == 8:\n            astrous_rates = [0, 2, 4, 8]\n        else:\n            raise Warning('Output stride must be 8 or 16!')\n\n        # astrous spational pyramid pooling part\n        self.conv1 = _ASPPConv(in_channels, out_channels, 1, 1)\n        self.conv2 = _ASPPConv(in_channels, out_channels, 3, 1, padding=astrous_rates[1], dilation=astrous_rates[1])\n        self.conv3 = _ASPPConv(in_channels, out_channels, 3, 1, padding=astrous_rates[2], dilation=astrous_rates[2])\n        self.conv4 = _ASPPConv(in_channels, out_channels, 3, 1, padding=astrous_rates[3], dilation=astrous_rates[3])\n\n        self.pool = nn.Sequential(\n            nn.AdaptiveAvgPool3d((1,1,1)),\n            nn.Conv3d(in_channels, out_channels, kernel_size=1, bias=False),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU()\n        )\n        self.bottleneck = nn.Sequential(\n            nn.Conv3d(out_channels * 5, out_channels, kernel_size=1, bias=False),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU()\n        )\n    \n    def forward(self, input):\n        input1 = self.conv1(input)\n        input2 = self.conv2(input)\n        input3 = self.conv3(input)\n        input4 = self.conv4(input)\n        \n        input5 = F.interpolate(self.pool(input), size=input4.size()[2:], mode='trilinear', align_corners=False)\n        output = torch.cat((input1, input2, input3, input4, input5), dim=1)\n\n        del input\n        del input1\n        del input2\n        del input3\n        del input4\n        del input5\n        torch.cuda.empty_cache()\n\n        output = self.bottleneck(output)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.048846Z","iopub.execute_input":"2022-09-23T09:36:29.049834Z","iopub.status.idle":"2022-09-23T09:36:29.068868Z","shell.execute_reply.started":"2022-09-23T09:36:29.04979Z","shell.execute_reply":"2022-09-23T09:36:29.067509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResUnet(nn.Module):\n\n    def __init__(self, in_ch, channels=16, blocks=3, use_aspp=False, is_aux=False):\n        super(ResUnet, self).__init__()\n\n        self.in_conv = DoubleConv(in_ch, channels, stride=2, kernel_size=3)\n        self.layer1 = make_res_layer(channels, channels * 2, blocks, stride=2)\n        self.layer2 = make_res_layer(channels * 2, channels * 4, blocks, stride=2)\n        self.layer3 = make_res_layer(channels * 4, channels * 8, blocks, stride=2)\n\n        self.up5 = nn.Upsample(scale_factor=2, mode='trilinear', align_corners=False)\n        self.conv5 = DoubleConv(channels * 12, channels * 4)\n        self.up6 = nn.Upsample(scale_factor=2, mode='trilinear', align_corners=False)\n        self.conv6 = DoubleConv(channels * 6, channels * 2)\n        self.up7 = nn.Upsample(scale_factor=2, mode='trilinear', align_corners=False)\n        self.conv7 = DoubleConv(channels * 3, channels)\n        self.up8 = nn.Upsample(scale_factor=2, mode='trilinear', align_corners=False)\n\n        self.aspp = ASPP(channels * 16, channels * 16)\n        self.is_aux = is_aux\n        self.use_aspp = use_aspp\n\n    def forward(self, input):\n        c1 = self.in_conv(input)\n        c2 = self.layer1(c1)\n        c3 = self.layer2(c2)\n        c4 = self.layer3(c3)\n\n        if self.use_aspp:\n            c4_ap = self.aspp(c4) \n        else:\n            c4_ap = c4\n\n        up_5 = self.up5(c4)\n        merge5 = torch.cat([up_5, c3], dim=1)\n        c5 = self.conv5(merge5)\n        up_6 = self.up6(c5)\n        merge6 = torch.cat([up_6, c2], dim=1)\n        c6 = self.conv6(merge6)\n        up_7 = self.up7(c6)\n        merge7 = torch.cat([up_7, c1], dim=1)\n        c7 = self.conv7(merge7)\n        up_8 = self.up8(c7)\n        if self.is_aux:\n            return [up_8, c7, c6, c5]\n        else:\n            return up_8","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.071189Z","iopub.execute_input":"2022-09-23T09:36:29.071808Z","iopub.status.idle":"2022-09-23T09:36:29.091652Z","shell.execute_reply.started":"2022-09-23T09:36:29.071761Z","shell.execute_reply":"2022-09-23T09:36:29.090144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegSigHead(nn.Module):\n\n    def __init__(self, in_channels, classes=1):\n        super(SegSigHead, self).__init__()\n        self.conv = nn.Conv3d(in_channels, 1, 1)\n        self.bce_loss_func = torch.nn.BCEWithLogitsLoss(reduce=False)\n\n    def forward(self, inputs):\n        features = self.conv(inputs)\n        return features\n\n    def forward_test(self, inputs):\n        return self.forward(inputs)\n\nclass SegSoftHead(nn.Module):\n\n    def __init__(self, in_channels, classes=14):\n        super(SegSoftHead, self).__init__()\n        self.conv = nn.Conv3d(in_channels, classes, 1)\n        self.multi_loss_func = torch.nn.CrossEntropyLoss(reduce=False)\n        self._classes = classes\n\n    def forward(self, inputs):\n        seg_predict = self.conv(inputs)\n        return seg_predict\n\n    def forward_test(self, inputs):\n        return self.forward(inputs)\n    \nclass SegSoftHead_cls_vert(nn.Module):\n\n    def __init__(self, in_channels, classes=5):\n        super(SegSoftHead, self).__init__()\n        self.conv = nn.Conv3d(in_channels, classes, 1)\n        self.bce_loss_func = torch.nn.BCEWithLogitsLoss(reduce=False)\n        self.multi_loss_func = torch.nn.CrossEntropyLoss(reduce=False)\n        self._classes = classes\n\n    def forward(self, inputs):\n        inputs = F.interpolate(inputs, scale_factor=1.0, mode='trilinear')\n        seg_predict = self.conv(inputs)\n\n        del inputs\n        torch.cuda.empty_cache()\n\n        return seg_predict\n\n    def forward_test(self, inputs):\n        return self.forward(inputs)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.094109Z","iopub.execute_input":"2022-09-23T09:36:29.094977Z","iopub.status.idle":"2022-09-23T09:36:29.113494Z","shell.execute_reply.started":"2022-09-23T09:36:29.094731Z","shell.execute_reply":"2022-09-23T09:36:29.111887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegNetwork(nn.Module):\n\n    def __init__(self, backbone,\n                head, apply_sync_batchnorm=False,\n                train_cfg=None,\n                test_cfg=None\n                ):\n        super(SegNetwork, self).__init__()\n        self.backbone = backbone\n        self.head = head\n        self._show_count = 0\n        if apply_sync_batchnorm:\n            self._apply_sync_batchnorm()\n\n    @torch.jit.ignore\n    def forward(self, vol, seg):\n        vol = vol.float()\n        seg = seg.float()\n        features = self.backbone(vol)\n        head_outs = self.head(features)\n        del features\n        del vol\n        del seg\n        torch.cuda.empty_cache()\n        return head_outs\n\n    @torch.jit.export\n    def forward_test(self, img):\n        features = self.backbone(img)\n        del img \n        torch.cuda.empty_cache()\n        seg_predict = self.head.forward_test(features)\n        del features\n        torch.cuda.empty_cache()\n        return seg_predict\n\n    def _apply_sync_batchnorm(self):\n        print('apply sync batch norm')\n        self.backbone = nn.SyncBatchNorm.convert_sync_batchnorm(self.backbone)\n        self.head = nn.SyncBatchNorm.convert_sync_batchnorm(self.head)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.118615Z","iopub.execute_input":"2022-09-23T09:36:29.119027Z","iopub.status.idle":"2022-09-23T09:36:29.133432Z","shell.execute_reply.started":"2022-09-23T09:36:29.118997Z","shell.execute_reply":"2022-09-23T09:36:29.13172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpineModel:\n\n    def __init__(self, model_f, network_f):\n        # TODO: 模型文件定制\n        self.model_f = model_f\n        self.network_f = network_f\n\nclass SpineConfig:\n\n    def __init__(self,  network_f):\n        # TODO: 模型配置文件\n        self.network_f = network_f\n        if self.network_f is not None:\n            if isinstance(self.network_f, str):\n                self.network_cfg = Config.fromfile(self.network_f)\n            else:\n                import tempfile\n\n                with tempfile.TemporaryDirectory() as temp_config_dir:\n                    with tempfile.NamedTemporaryFile(dir=temp_config_dir, suffix='.py') as temp_config_file:\n                        with open(temp_config_file.name, 'wb') as f:\n                            f.write(self.network_f.read())\n\n                        self.network_cfg = Config.fromfile(temp_config_file.name)\n\n    def __repr__(self) -> str:\n        return str(self.__dict__)\n\nclass ClsPredictor2D:\n\n    def __init__(self, gpu: int, model: SpineModel):\n        self.gpu = gpu\n        self.model = model\n        self.config = SpineConfig(self.model.network_f)\n        self.load_model()\n    def load_model(self):\n        self.net = self._load_model(self.model.model_f, self.config.network_cfg, half=False)\n\n    def _load_model(self, model_f, network_f, half=False) -> None:\n        net = self.load_model_pth(model_f, network_f, half)\n        return net\n\n    def load_model_pth(self, model_f, network_cfg, half) -> None:\n        # 加载动态图\n        config = network_cfg\n        backbone = ResNet2D_50()\n        head = ClsHead_Res(in_channels=config.model[\"in_channels\"], num_classes=config.model[\"num_classes\"])\n        net = ClsNetwork_Res(backbone, head)\n        checkpoint = torch.load(model_f, map_location=f\"cpu\")\n        if list(checkpoint[\"state_dict\"].keys())[0].startswith(\"module.\"):\n            state_dict = {k[7:]: v for k, v in checkpoint[\"state_dict\"].items()}\n        else:\n            state_dict = checkpoint[\"state_dict\"]\n        net.load_state_dict(state_dict, strict=False)\n        net.eval()\n        net.half()\n        net.cuda()\n        net = net.forward_test\n        return net\n    \n    def _get_input(self, img, spacing_zyx):\n\n        config = self.config.network_cfg\n\n\n        def _window_array(img, win_level, win_width):\n            win = [\n                win_level - win_width / 2,\n                win_level + win_width / 2,\n            ]\n            img = torch.clamp(img, win[0], win[1])\n            img -= win[0]\n            img /= win_width\n            return img\n\n        patch_size = np.array(config.patch_size)\n        all_size = np.array(list(img.shape[:2]) + list(patch_size))\n        if np.any(all_size != img.shape):\n            img = zoom(img, np.array(np.array(all_size) / img.shape), order=1)\n        img = torch.from_numpy(img).float()\n        img = [_window_array(img, wl, wd) for wl, wd in zip(config.win_level, config.win_width)]\n        img = torch.cat(img, dim=1)\n        img = img.detach()\n        return img\n\n    def forward(self, img, spacing):\n        # img: batch_size, channel, width, length, ...\n        spacing_zyx = np.array(spacing)\n        data = self._get_input(img, spacing_zyx)\n\n        with autocast():\n            data = data.cuda().detach()\n            pred = self.net(data)\n            # print(pred.size())\n            del data\n            if pred.size()[1] > 1:\n                pred = F.softmax(pred, dim=1)\n            else:\n                pred = torch.sigmoid(pred)\n        pred = pred.cpu().detach().numpy()\n        return pred\n\n\nclass ClsPredictor2D_local(ClsPredictor2D):\n\n    def load_model_pth(self, model_f, network_cfg, half) -> None:\n        # 加载动态图\n        config = network_cfg\n        backbone = ResNet2D_50()\n        head = ClsHead_Res(in_channels=config.model[\"in_channels\"], num_classes=config.model[\"num_classes\"])\n        net = ClsNetwork_Res(backbone, head)\n        checkpoint = torch.load(model_f, map_location=f\"cpu\")\n        if list(checkpoint[\"state_dict\"].keys())[0].startswith(\"module.\"):\n            state_dict = {k[7:]: v for k, v in checkpoint[\"state_dict\"].items()}\n        else:\n            state_dict = checkpoint[\"state_dict\"]\n        net.load_state_dict(state_dict, strict=False)\n        net.eval()\n        net.half()\n        net.cuda()\n        net = net.forward_test\n        return net","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.136146Z","iopub.execute_input":"2022-09-23T09:36:29.136754Z","iopub.status.idle":"2022-09-23T09:36:29.16885Z","shell.execute_reply.started":"2022-09-23T09:36:29.136708Z","shell.execute_reply":"2022-09-23T09:36:29.167328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass ClsHead_Res(nn.Module):\n\n    def __init__(self, in_channels, num_classes):\n        super(ClsHead_Res, self).__init__()\n        self.fc = nn.Linear(in_channels, num_classes)\n        self.v = nn.NLLLoss()\n        self.classes = num_classes\n    def forward(self, inputs):\n        cls_out = torch.flatten(inputs, 1)        \n        cls_out = self.fc(cls_out)\n        return cls_out\n\n    def forward_test(self, inputs):\n        return self.forward(inputs)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.172212Z","iopub.execute_input":"2022-09-23T09:36:29.174213Z","iopub.status.idle":"2022-09-23T09:36:29.18564Z","shell.execute_reply.started":"2022-09-23T09:36:29.174116Z","shell.execute_reply":"2022-09-23T09:36:29.183885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClsNetwork_Res(nn.Module):\n\n    def __init__(self, backbone,\n                head, apply_sync_batchnorm=False,\n                train=True,\n                train_cfg=None,\n                test_cfg=None\n                ):\n        super(ClsNetwork_Res, self).__init__()\n        self.backbone = backbone\n        self.head = head\n        self._show_count = 0\n        self._train = train\n        if apply_sync_batchnorm:\n            self._apply_sync_batchnorm()\n\n    @torch.jit.ignore\n    def forward(self, img, label):\n        with torch.no_grad():\n            img = img.float()\n            # print('shape of image, ', img.shape)\n            label = label.float()\n        features = self.backbone(img)\n        head_outs = self.head(features)\n        return head_outs\n\n    @torch.jit.export\n    def forward_test(self, img):\n        features = self.backbone(img)\n        label_predict = self.head(features)\n        return label_predict\n\n    def single_test(self, img, label):\n        with torch.no_grad():\n            img = img.float()\n            features = self.backbone(img)\n        predict = self.head(features)\n        # predict = torch.sigmoid(predict)\n        predict = F.softmax(predict, dim=1)\n        predict = torch.argmax(predict, dim=1)\n        acc_value = []\n        label = torch.squeeze(label)\n        labels = torch.argmax(label, dim=1)\n        for i in range(labels.shape[0]):\n            if (predict[i]==labels[i]):\n                acc_value.append([1])\n            else:\n                acc_value.append([0])\n        \n        return acc_value\n\n    def _apply_sync_batchnorm(self):\n        print('apply sync batch norm')\n        self.backbone = nn.SyncBatchNorm.convert_sync_batchnorm(self.backbone)\n        self.head = nn.SyncBatchNorm.convert_sync_batchnorm(self.head)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.188032Z","iopub.execute_input":"2022-09-23T09:36:29.188963Z","iopub.status.idle":"2022-09-23T09:36:29.206287Z","shell.execute_reply.started":"2022-09-23T09:36:29.188883Z","shell.execute_reply":"2022-09-23T09:36:29.204312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 导入相关的模块\nimport torch\nimport torch.nn as nn\n \n# 本文提供的所有的Resnet类型\n__all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101',\n           'resnet152', 'resnext50_32x4d', 'resnext101_32x8d',\n           'wide_resnet50_2', 'wide_resnet101_2']\n# 3x3卷积\n# in_planes是输入图像的channel，out_planes是输出图像的channel\ndef conv3x3_2d(in_planes, out_planes, stride=1, groups=1, dilation=1):\n    \"\"\"3x3 convolution with padding\"\"\"\n    return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,\n                     padding=dilation, groups=groups, bias=False, dilation=dilation)\n \n# 1x1卷积 \ndef conv1x1_2d(in_planes, out_planes, stride=1):\n    \"\"\"1x1 convolution\"\"\"\n    return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)   \n\n# resnet18和34的block\nclass BasicBlock2D(nn.Module):\n    expansion = 1\n    __constants__ = ['downsample']\n \n    def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1,\n                 base_width=64, dilation=1, norm_layer=None):\n        super(BasicBlock2D, self).__init__()\n        if norm_layer is None:\n            norm_layer = nn.BatchNorm2d # Batch Normalization\n        if groups != 1 or base_width != 64:\n            raise ValueError('BasicBlock2D only supports groups=1 and base_width=64')\n        if dilation > 1:\n            raise NotImplementedError(\"Dilation > 1 not supported in BasicBlock2D\")\n        # Both self.conv1 and self.downsample layers downsample the input when stride != 1\n        self.conv1 = conv3x3_2d(inplanes, planes, stride)\n        self.bn1 = norm_layer(planes)\n        self.relu = nn.ReLU(inplace=True) # Relu激活函数\n        self.conv2 = conv3x3_2d(planes, planes)\n        self.bn2 = norm_layer(planes)\n        self.downsample = downsample\n        self.stride = stride\n# 下面是block的基本结构，每一层是如何组成的，都很清楚\n    def forward(self, x):\n        identity = x\n \n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n \n        out = self.conv2(out)\n        out = self.bn2(out)\n \n        if self.downsample is not None:\n            identity = self.downsample(x)\n \n        out += identity\n        out = self.relu(out)\n \n        return out\n    \nclass Bottleneck2D(nn.Module):\n    expansion = 4\n    __constants__ = ['downsample']\n \n    def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1,\n                 base_width=64, dilation=1, norm_layer=None):\n        super(Bottleneck2D, self).__init__()\n        if norm_layer is None:\n            norm_layer = nn.BatchNorm2d\n        width = int(planes * (base_width / 64.)) * groups\n        # Both self.conv2 and self.downsample layers downsample the input when stride != 1\n        self.conv1 = conv1x1_2d(inplanes, width)\n        self.bn1 = norm_layer(width)\n        self.conv2 = conv3x3_2d(width, width, stride, groups, dilation)\n        self.bn2 = norm_layer(width)\n        # 输出的通道数由64变为256,所以下面要乘self.expansion\n        self.conv3 = conv1x1_2d(width, planes * self.expansion)\n        self.bn3 = norm_layer(planes * self.expansion)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stride = stride\n \n    def forward(self, x):\n        identity = x\n \n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n \n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n \n        out = self.conv3(out)\n        out = self.bn3(out)\n \n        if self.downsample is not None:\n            identity = self.downsample(x)\n \n        out += identity\n        out = self.relu(out)\n \n        return out\n    \nclass ResNet(nn.Module):\n     \n    def __init__(self, block, layers, num_classes=1000, zero_init_residual=False,\n                 groups=1, width_per_group=64, replace_stride_with_dilation=None,\n                 norm_layer=None):\n        super(ResNet, self).__init__()\n        if norm_layer is None:\n            norm_layer = nn.BatchNorm2d\n        self._norm_layer = norm_layer\n \n        self.inplanes = 64\n        self.dilation = 1\n        if replace_stride_with_dilation is None:\n            # each element in the tuple indicates if we should replace\n            # the 2x2 stride with a dilated convolution instead\n            replace_stride_with_dilation = [False, False, False]\n        if len(replace_stride_with_dilation) != 3:\n            raise ValueError(\"replace_stride_with_dilation should be None \"\n                             \"or a 3-element tuple, got {}\".format(replace_stride_with_dilation))\n        self.groups = groups\n        self.base_width = width_per_group\n        # 由图5可知，网络第一层是7x7的卷积层，stride为2，输入3通道，输出64通道\n        self.conv1 = nn.Conv2d(1, self.inplanes, kernel_size=7, stride=2, padding=3,\n                               bias=False)\n        self.bn1 = norm_layer(self.inplanes)\n        self.relu = nn.ReLU(inplace=True)\n        # 由图5可知，网络第二层是3x3的max pooling层，stride为2\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        # 对于resnet18，layers=[2,2,2,2]，也就是相同的block重复的次数\n        self.layer1 = self._make_layer(block, 64, layers[0])\n        self.layer2 = self._make_layer(block, 128, layers[1], stride=2,\n                                       dilate=replace_stride_with_dilation[0])\n        self.layer3 = self._make_layer(block, 256, layers[2], stride=2,\n                                       dilate=replace_stride_with_dilation[1])\n        self.layer4 = self._make_layer(block, 512, layers[3], stride=2,\n                                       dilate=replace_stride_with_dilation[2])\n        # 把特征图变为1x1大小\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        # 全连接层，将输出的通道数由512*block.expansion变为num_calsses\n        self.fc = nn.Linear(512 * block.expansion, num_classes)\n \n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n \n        # Zero-initialize the last BN in each residual branch,\n        # so that the residual branch starts with zeros, and each residual block behaves like an identity.\n        # This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677\n        if zero_init_residual:\n            for m in self.modules():\n                if isinstance(m, Bottleneck2D):\n                    nn.init.constant_(m.bn3.weight, 0)\n                elif isinstance(m, BasicBlock2D):\n                    nn.init.constant_(m.bn2.weight, 0)\n \n    def _make_layer(self, block, planes, blocks, stride=1, dilate=False):\n        norm_layer = self._norm_layer\n        downsample = None\n        previous_dilation = self.dilation\n        if dilate:\n            self.dilation *= stride\n            stride = 1\n        ''' layer1的时候不执行下面这个if，layer2往后，由于self.inplanes不等于\n        planes*block.expansion，所以需要经过下面的if段，作用是将通道数由self.planes变为\n        planes*block.expansion'''\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            downsample = nn.Sequential(\n                conv1x1_2d(self.inplanes, planes * block.expansion, stride),\n                norm_layer(planes * block.expansion),\n            )\n \n        layers = []\n        layers.append(block(self.inplanes, planes, stride, downsample, self.groups,\n                            self.base_width, previous_dilation, norm_layer))\n        self.inplanes = planes * block.expansion\n        # range取[1,blocks)的值（前闭后开能懂吧）\n        for _ in range(1, blocks):\n            layers.append(block(self.inplanes, planes, groups=self.groups,\n                                base_width=self.base_width, dilation=self.dilation,\n                                norm_layer=norm_layer))\n        # 经过上面的操作，相同的block重复了blocks-1次\n \n        return nn.Sequential(*layers)\n \n    def _forward_impl(self, x):\n        # See note [TorchScript super()]\n        # 真正的结构在这里哈\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n \n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n \n        x = self.avgpool(x)\n        # 把2、3、4维的值相乘变为1维，所以最后x变为2维[batch_size,size]\n        x = torch.flatten(x, 1)\n        # 全连接层的输入输出一般为2维张量，这也解释了上一步为什么要这么做\n        x = self.fc(x)\n \n        return x\n \n    def forward(self, x):\n        return self._forward_impl(x)\n\n\nclass ResNet2D_18(nn.Module):\n\n    def __init__(self, **kwargs):\n        super(ResNet2D_18, self).__init__()\n        self.feature_extractor = ResNet(BasicBlock2D, [2, 2, 2, 2], **kwargs)\n\n    def forward(self, input):\n        return self.feature_extractor(input)\n\nclass ResNet2D_34(nn.Module):\n\n    def __init__(self, **kwargs):\n        super(ResNet2D_34, self).__init__()\n        self.feature_extractor = ResNet(BasicBlock2D, [3, 4, 6, 3], **kwargs)\n\n    def forward(self, input):\n        return self.feature_extractor(input)\n\nclass ResNet2D_50(nn.Module):\n\n    def __init__(self, **kwargs):\n        super(ResNet2D_50, self).__init__()\n        self.feature_extractor = ResNet(Bottleneck2D, [3, 4, 6, 3], **kwargs)\n\n    def forward(self, input):\n        return self.feature_extractor(input)\n\nclass ResNet2D_101(nn.Module):\n\n    def __init__(self, **kwargs):\n        super(ResNet2D_101, self).__init__()\n        self.feature_extractor = ResNet(Bottleneck2D, [3, 4, 23, 3], **kwargs)\n\n    def forward(self, input):\n        return self.feature_extractor(input)","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.2087Z","iopub.execute_input":"2022-09-23T09:36:29.209438Z","iopub.status.idle":"2022-09-23T09:36:29.259564Z","shell.execute_reply.started":"2022-09-23T09:36:29.209394Z","shell.execute_reply":"2022-09-23T09:36:29.258121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference(predictor: SegPredictor, hu_volume, spacing):\n    pred_array = predictor.forward(hu_volume, spacing)\n    return pred_array\n\ndef crop_patch(vol, z_num=10):\n    z, x, y = vol.shape\n    if z%z_num:\n        z_stride = int(z/z_num) + 1\n    else:\n        z_stride = int(z/z_num)\n    vol_crop = []\n    z_max_list = []\n    z_min_list = []\n\n    for z_item in range(z_num):\n        z_min = z_item*z_stride\n        z_max = (z_item+1)*z_stride\n        if z_max > z:\n            z_max = z\n            z_min = z-z_stride\n\n        vol_item = vol[z_min:z_max]\n        vol_crop.append(vol_item)\n        z_max_list.append(z_max)\n        z_min_list.append(z_min)\n    return vol_crop, z_max_list, z_min_list\n\ndef find_valid_region(mask, values, low_margin=[0,0,0], up_margin=[0,0,0]):\n    for v in values:\n        mask[mask == v] = 100\n    nonzero_points = np.argwhere((mask > 20))\n    if len(nonzero_points) == 0:\n        return None, None\n    else:\n        v_min = np.min(nonzero_points, axis=0)\n        v_max = np.max(nonzero_points, axis=0)\n        assert len(v_min) == len(low_margin), f'the length of margin is not equal the mask dims {len(v_min)}!'\n        for idx in range(len(v_min)):\n            v_min[idx] = max(0, v_min[idx] - low_margin[idx])\n            v_max[idx] = min(mask.shape[idx], v_max[idx] + up_margin[idx])\n        return v_min, v_max","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.262302Z","iopub.execute_input":"2022-09-23T09:36:29.264939Z","iopub.status.idle":"2022-09-23T09:36:29.281155Z","shell.execute_reply.started":"2022-09-23T09:36:29.264871Z","shell.execute_reply":"2022-09-23T09:36:29.279704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    import cc3d\nexcept:\n    !cp ../input/mirror-cc3d/connected_components_3d-3.10.2-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl ./\n    !pip install ./connected_components_3d-3.10.2-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n    import cc3d\n\ndef max_connected_region(mask, values):\n    out = np.zeros(mask.shape, dtype=np.uint8)\n    for v in values:\n        temp = mask.copy()\n        temp[temp != v] = 0\n        if v in list(np.unique(temp)):\n            labeled, N = cc3d.connected_components((temp > 0), return_N=True)\n            area = np.sum(labeled == 1); target = 1\n            if N >= 2:\n                for idx in range(2, N+1):\n                    if np.sum(labeled == idx) > area:\n                        target = idx\n                        area = np.sum(labeled == idx)\n            assert (target > 0)\n            out[labeled == target] = v\n    return out\n\ndef main(input_path, output_path, gpu):\n    \n    model_seg_global = '../input/test-model/seg_model/coarse_epoch_60.pth'\n    network_seg_global = '../input/test-model/seg_model/train_config_coarse.py'\n    \n#     model_seg_fine = '../input/model-0919/fine_eight_epoch_90.pth'\n#     network_seg_fine = '../input/model-0919/train_config_fine_eight.py'\n    \n    model_seg_fine = '../input/model-0919/fine_epoch_30.pth'\n    network_seg_fine = '../input/model-0919/train_config_fine.py'\n    \n#     model_seg_fine = '../input/test-model/seg_model/fine_epoch_50.pth'\n#     network_seg_fine = '../input/test-model/seg_model/train_config_fine.py'\n    \n    model_seg_crop = '../input/test-model/seg_model/crop_epoch_30.pth'\n    network_seg_crop = '../input/test-model/seg_model/train_config_crop.py'\n    \n    model_cls_vert = '../input/test-model/cls_epoch_120_16.pth'\n    network_cls_vert = '../input/test-model/train_config_cls_16.py'\n    \n    model_cls_frac = '../input/model-0919/frac_2D_epoch_43.pth'\n    network_cls_frac = '../input/model-0919/train_config_res_frac.py'\n    \n    model_segmask_global = SegModel(\n        model_f=model_seg_global,\n        network_f=network_seg_global,\n    )\n    model_segmask_fine = SegModel(\n        model_f=model_seg_fine,\n        network_f=network_seg_fine,\n    )\n    model_segmask_crop = SegModel(\n        model_f=model_seg_crop,\n        network_f=network_seg_crop,\n    )\n    model_seg_cls = SegModel(\n        model_f=model_cls_vert,\n        network_f=network_cls_vert,\n    )\n    model_frac_cls = SpineModel(\n        model_f = model_cls_frac,\n        network_f = network_cls_frac\n    )\n    \n    predictor_segmask_global = SegPredictor_(\n        gpu = gpu,\n        model = model_segmask_global,\n    )\n    predictor_segmask_fine = SegPredictor_(\n        gpu = gpu,\n        model = model_segmask_fine,\n    )\n    predictor_segmask_crop = SegPredictor_(\n        gpu = gpu,\n        model = model_segmask_crop,\n    )\n    predictor_seg_cls = SegPredictor(\n        gpu = gpu,\n        model = model_seg_cls,\n    )\n    predictor_frac_cls = ClsPredictor2D_local(\n        gpu = gpu,\n        model = model_frac_cls,\n    )\n\n\n    os.makedirs(output_path, exist_ok=True)\n    file_result = open('submission.csv', 'w')\n    writer = csv.writer(file_result)\n#     row_head = ['row_id', 'patient_level', 'c1', 'c2', 'c3', 'c4', 'c5', 'c6', 'c7']\n    row_head = ['StudyInstanceUID','row_id', 'fractured']\n    writer.writerow(row_head)\n    \n    pids = sorted(os.listdir(input_path))\n    \n    predictions = []\n    batch_size_2D = 256\n    batch_size_3D = 2\n    \n    for pid in tqdm(pids):\n        sub_res = [pid]\n        pid_path = os.path.join(input_path, pid)\n        if not os.path.isdir(pid_path):\n            continue\n        dcms = sorted(os.listdir(pid_path))\n        if not dcms[0].endswith('.dcm'):\n            continue\n        print('data load ......') \n        hu_volume, spacing = dcm2nii(pid_path)\n#         hu_volume = sitk.GetArrayFromImage(sitk_vol)\n        print(pid, hu_volume.shape)\n        src_shape = hu_volume.shape\n#         spacing = sitk_vol.GetSpacing()\n        spacing = spacing[::-1]\n\n        global_seg = inference(predictor_segmask_global, hu_volume, spacing)\n        global_seg_shape = global_seg.shape\n        heatmap = global_seg.copy()\n        \n        # fine seg\n        pmin, pmax = find_valid_region(heatmap.copy(), [1])\n        vol_case = hu_volume[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] \n        pred_case = inference(predictor_segmask_fine, vol_case, spacing) \n        # print(np.unique(pred_case))\n        heatmap[heatmap > 0] = 0\n        heatmap[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] = pred_case\n        \n#         # crop fine seg #\n#         temp_ori = np.zeros(heatmap.shape, dtype=np.uint8)\n#         pmin, pmax = find_valid_region(global_seg.copy(), [1], low_margin=[0, 0, 0], up_margin=[0, 0, 0])\n#         vol_case = hu_volume[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] \n#         temp_ = np.zeros(vol_case.shape, dtype=np.uint8)\n#         crop_vol, z_max_list, z_min_list = crop_patch(vol_case)\n#         for i in range(len(crop_vol)):\n#             # print(crop_vol[i].shape)\n#             pred_case_item = inference(predictor_segmask_crop, crop_vol[i], spacing) \n#             # print(z_min_list[i], z_max_list[i], temp_.shape, pred_case_item.shape) \n#             temp_[z_min_list[i]:z_max_list[i]] = pred_case_item\n#         temp_ori[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] = temp_\n#         heatmap[temp_ori == 1] = 1; heatmap[temp_ori == 0] = 0\n        \n        # color\n        pmin, pmax = find_valid_region(heatmap.copy(), [1])\n        if pmin is None or pmax is None:\n            print(\"Wrong spine seg mask!\")\n        else:\n            seg_case = heatmap[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]]\n        seg_cls =  inference(predictor_seg_cls, seg_case, spacing)\n        heatmap = np.zeros(src_shape, dtype=np.uint8)\n        heatmap[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] = seg_cls.astype(np.uint8)\n\n        # for frac classification\n#         vert_labels = sorted(list(np.unique(heatmap)))\n#         vert_labels.remove(0)\n#         heatmap = max_connected_region(heatmap, vert_labels)\n        \n        vert_labels = sorted(list(np.unique(heatmap)))\n        vert_labels.remove(0)\n        frac_probs = []\n        patch_size = (224, 224)\n        slices2D = []\n        for vert_num in vert_labels:\n            pmin, pmax = find_valid_region(heatmap.copy(), [vert_num])\n            vol_case = hu_volume[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] \n            target_shape = np.array(list(vol_case.shape[:1]) + list(patch_size))\n            vol_case = zoom(vol_case, np.array(target_shape / vol_case.shape), order=1)\n            slices2D.append([vol_case, [vert_num] * (pmax[0]-pmin[0])])\n        vol_case = [s[0] for s in slices2D]\n        C_ID_case = [s[1] for s in slices2D]\n        vol_case = np.concatenate(vol_case, axis=0)\n        C_ID_case = np.concatenate(C_ID_case, axis=0)\n        assert vol_case.shape[0] == C_ID_case.shape[0], \"2D slices num and C indexes dismatch!\"\n        idx = 0\n        while idx < vol_case.shape[0]:\n            slices2D = vol_case[idx:idx+batch_size_2D]\n            slices2D = slices2D[:, None, :, :]\n            \n            pred_case = inference(predictor_frac_cls, slices2D, spacing)[:, 1]\n#             pred_case = inference(predictor_frac_cls_local, slices2D, spacing)[:, 1]\n            frac_probs.extend(list(pred_case))\n            idx += batch_size_2D\n        frac_probs = np.array(frac_probs)\n        C_frac_prob = []\n        assert frac_probs.shape == C_ID_case.shape, \"2D slices num and C indexes dismatch!\"\n        for vert_num in vert_labels:\n            frac_vert_prob = sorted(frac_probs[C_ID_case == vert_num])\n            frac_vert = max(1e-7, min(np.mean(frac_vert_prob), 0.9995))\n            predictions.append([pid, pid+'_C'+str(vert_num), frac_vert])\n            C_frac_prob.append(frac_vert)\n            print('fracture prob of vert: ', vert_num, frac_vert)\n        patient_overall = patient_overall = max(min(0.99, 1 - np.prod(1 - np.array(C_frac_prob))), 1e-7)\n        predictions.append([pid, pid+'_patient_overall', patient_overall])\n        \n        \n\n    data_res = pd.DataFrame(\n        predictions, columns=[\"StudyInstanceUID\", \"row_id\", \"fractured\"]\n    )\n    data_res.to_csv('predictions.csv', encoding='utf8', index=None)\n    data_res[['row_id', 'fractured']].to_csv('submission_.csv', encoding='utf8', index=None)\n    print(data_res[['row_id', 'fractured']])\n    !mv 'submission_.csv' 'submission.csv'\n#         temp_ori = np.zeros(heatmap.shape, dtype=np.uint8)\n#         pmin, pmax = find_valid_region(global_seg.copy(), [1], low_margin=[0, 0, 0], up_margin=[0, 0, 0])\n#         vol_case = hu_volume[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] \n#         temp_ = np.zeros(vol_case.shape, dtype=np.uint8)\n#         crop_vol, z_max_list, z_min_list = crop_patch(vol_case)\n\n#         for i in range(len(crop_vol)):\n#             # print(crop_vol[i].shape)\n#             pred_case_item = inference(predictor_segmask_crop, crop_vol[i], spacing) \n#             # print(z_min_list[i], z_max_list[i], temp_.shape, pred_case_item.shape) \n#             temp_[z_min_list[i]:z_max_list[i]] = pred_case_item\n#         temp_ori[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] = temp_\n#         heatmap[temp_ori == 1] = 1; heatmap[temp_ori == 0] = 0\n        \n#         pmin, pmax = find_valid_region(global_seg.copy(), [1])\n#         if pmin is None or pmax is None:\n#             print(\"Wrong spine seg mask!\")\n#         else:\n#             seg_case = heatmap[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]]\n#         seg_cls =  inference(predictor_seg_cls, seg_case, spacing)\n#         heatmap = np.zeros(src_shape, dtype=np.uint8)\n#         heatmap[pmin[0]:pmax[0], pmin[1]:pmax[1], pmin[2]:pmax[2]] = seg_cls.astype(np.uint8)\n        \n#         frac_prob = []\n\n#         for vert_num in range(1, np.max(np.unique(heatmap))+1):\n#             pmin, pmax = find_valid_region(heatmap.copy(), [vert_num])\n#             frac_vert_prob = []\n#             for j in range(pmin[0], pmax[0]):\n#                 pred_case_item = inference(predictor_frac_cls, hu_volume[j], spacing)\n#                 pred_case_item_prob = pred_case_item[0][1]\n#                 frac_vert_prob.append(pred_case_item_prob)\n#             frac_vert_prob_array = np.array(frac_vert_prob)\n# #             print('fracture prob of vert: ', vert_num, np.max(frac_vert_prob_array))\n#             sub_res.append(np.max(frac_vert_prob_array))\n#             frac_prob.append(np.max(frac_vert_prob_array))\n#         for frac_num in range(len(frac_prob)):\n        \n#             frac_num_write = pid+'_C'+str(frac_num+1)\n#             fractured_write = frac_prob[frac_num]\n#             sub_res = [pid,frac_num_write, fractured_write]\n#             writer.writerow(sub_res)\n        \n\n\nif __name__ == '__main__':\n    input_path = '../input/rsna-2022-cervical-spine-fracture-detection/test_images'\n    output_path = './'\n    main(\n        input_path=input_path,\n        output_path=output_path,\n        gpu=0,\n    )","metadata":{"execution":{"iopub.status.busy":"2022-09-23T09:36:29.285885Z","iopub.execute_input":"2022-09-23T09:36:29.286321Z","iopub.status.idle":"2022-09-23T09:37:38.668812Z","shell.execute_reply.started":"2022-09-23T09:36:29.286293Z","shell.execute_reply":"2022-09-23T09:37:38.667034Z"},"trusted":true},"execution_count":null,"outputs":[]}]}