| 98 | return grad |
| 99 | |
| 100 | class BaseAnalysis(BaseModel): |
| 101 | def __init__(self, opt): |
| 102 | self.opt = opt |
| 103 | self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') |
| 104 | self.setup_environment() |
| 105 | self._load_model() |
| 106 | self.hook_list = self._register_hooks() |
| 107 | self.test_loaders = self.build_test_loaders() |
| 108 | |
| 109 | def setup_environment(self): |
| 110 | torch.backends.cudnn.benchmark = True |
| 111 | |
| 112 | # mkdir and initialize loggers |
| 113 | make_exp_dirs(self.opt) |
| 114 | log_file = osp.join(self.opt['path']['log'], f"test_{self.opt['name']}_{get_time_str()}.log") |
| 115 | self.logger = get_root_logger(logger_name='ram', log_level=logging.INFO, log_file=log_file) |
| 116 | self.logger.info(get_env_info()) |
| 117 | self.logger.info(dict2str(self.opt)) |
| 118 | |
| 119 | |
| 120 | def _load_model(self): |
| 121 | self.model = build_network(self.opt['network_g'], False).to(self.device) |
| 122 | self.model.eval() |
| 123 | self.model = self.model.to(self.device) |
| 124 | load_path = self.opt['path'].get('pretrain_network_g') |
| 125 | if load_path: |
| 126 | self.load_network(self.model, load_path, strict=False) |
| 127 | |
| 128 | def _register_hooks(self): |
| 129 | module_name_list = [] |
| 130 | hook_list = [] |
| 131 | name_list = [] |
| 132 | def get_module_from_name(name): |
| 133 | name_parts = name.split('.')[:-1] |
| 134 | module_name = 'self.model' |
| 135 | |
| 136 | for part in name_parts: |
| 137 | if part == 'mask_token' or part == 'weight' or part == 'bias': |
| 138 | continue |
| 139 | if part.isdigit(): |
| 140 | module_name += f'[{part}]' |
| 141 | else: |
| 142 | module_name += f'.{part}' |
| 143 | return module_name, '.'.join(name_parts) |
| 144 | |
| 145 | for name,param in self.model.named_parameters(): |
| 146 | module_name,name = get_module_from_name(name) |
| 147 | |
| 148 | # print(module_name) |
| 149 | if module_name != 'self.model': |
| 150 | if len(module_name_list)==0 or module_name_list[-1] != module_name: |
| 151 | module_name_list.append(module_name) |
| 152 | name_list.append(name) |
| 153 | module = eval(module_name) |
| 154 | hook_list.append(Hook_back_loop(module, name)) |
| 155 | |
| 156 | return hook_list |
| 157 |
nothing calls this directly
no outgoing calls
no test coverage detected