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

Class BaseAnalysis

scripts/analysis_utils.py:100–195  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

98 return grad
99
100class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected