| 8 | from ram.archs.AdaSAM_arch import AdaptiveMaskPixGenerator |
| 9 | |
| 10 | class Ada_MACAnalysis(BaseAnalysis): |
| 11 | def __init__(self, opt): |
| 12 | super().__init__(opt) |
| 13 | self.mask_generator = self._load_mask_generator(opt) |
| 14 | self.mask_generator.eval() |
| 15 | |
| 16 | def _load_mask_generator(self, opt): |
| 17 | mask_opt = opt.get('network_mask', {}) |
| 18 | mask_generator = build_network(mask_opt).to(self.device) |
| 19 | mask_path_opt = opt.get('net_mask_path', {}) |
| 20 | mask_path = mask_path_opt.get('path', None) |
| 21 | self.load_network( |
| 22 | mask_generator, |
| 23 | mask_path, |
| 24 | strict=mask_path_opt.get('strict_load', True), |
| 25 | param_key=mask_path_opt.get('param_key', 'params') |
| 26 | ) |
| 27 | mask_generator.eval() |
| 28 | return mask_generator |
| 29 | |
| 30 | def analyze(self): |
| 31 | total_filter_mac = [0.0] * len(self.hook_list) |
| 32 | for test_loader in self.test_loaders: |
| 33 | test_set_name = test_loader.dataset.opt['name'] |
| 34 | num_samples = self.opt.get('num_samples',10) |
| 35 | print(f'Analyzing {test_set_name}..\n') |
| 36 | pbar = tqdm(total=num_samples, desc='') |
| 37 | for idx, val_data in enumerate(test_loader): |
| 38 | if idx >= num_samples: |
| 39 | break |
| 40 | tensor_lq = val_data['lq'].to(self.device) |
| 41 | imgname = osp.basename(val_data['lq_path'][0]) |
| 42 | tensor_base = torch.zeros_like(tensor_lq) |
| 43 | layer_conductance = self._mask_attribute_conductance(tensor_base, tensor_lq) |
| 44 | total_filter_mac = [a + b for a, b in zip(total_filter_mac, layer_conductance)] |
| 45 | pbar.set_description(f'Read {imgname}') |
| 46 | pbar.update(1) |
| 47 | self._save_results(total_filter_mac, 'mac') |
| 48 | |
| 49 | def _mask_attribute_conductance(self, base_img, final_img): |
| 50 | total_step = self.opt['total_step'] |
| 51 | |
| 52 | with torch.no_grad(): |
| 53 | _, _, p_x_full = self.mask_generator(final_img) |
| 54 | p_x_flat = p_x_full.flatten() |
| 55 | order_array = torch.argsort(p_x_flat).cpu().numpy() |
| 56 | |
| 57 | start_ratio = self.opt['pretrained_ratio'] |
| 58 | all_hook_layer_conductance = [0.0] * len(self.hook_list) |
| 59 | last_hook_layer_output = [] |
| 60 | |
| 61 | for step in range(total_step): |
| 62 | alpha = 1 - start_ratio + start_ratio * step / total_step |
| 63 | interpolated_img = self._get_interpolated_img_from_mask_attribute_path(base_img, final_img, alpha, order_array).to(self.device) |
| 64 | self.model.zero_grad() |
| 65 | interpolated_output = self.model(interpolated_img,None,None) |
| 66 | |
| 67 | if isinstance(interpolated_output, tuple): |