MCPcopy Create free account
hub / github.com/DragonisCV/RAM / Ada_MACAnalysis

Class Ada_MACAnalysis

scripts/adaSAM_mac_analysis.py:10–95  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8from ram.archs.AdaSAM_arch import AdaptiveMaskPixGenerator
9
10class 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):

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected