• OpenMMLap之Hook机制详解


    HOOK
    HOOK机制在OpenMMLab系列框架中应用广泛,结合Runner类可以实现训练过程中的整个生命周期的管理。例如调整学习率,保存模型,优化器等
    通过register的形式诸如Runner中实现丰富的扩展功能。接下来我们以train工作流为例分析调用位置及机制。
    1.调用位置
    以EpochBasedRunner(BaseRunner)为例分析

    mmcv/runner/epoch_base_runner.py

    1. class EpochBasedRunner(BaseRunner):
    2. """Epoch-based Runner.
    3. This runner train models epoch by epoch.
    4. """
    5. def train(self, data_loader, **kwargs):
    6. self.model.train()
    7. self.mode = 'train'
    8. self.data_loader = data_loader
    9. self._max_iters = self._max_epochs * len(self.data_loader)
    10. ######1.
    11. self.call_hook('before_train_epoch')
    12. time.sleep(2) # Prevent possible deadlock during epoch transition
    13. for i, data_batch in enumerate(self.data_loader):
    14. self.data_batch = data_batch
    15. self._inner_iter = i
    16. ######2.
    17. self.call_hook('before_train_iter')
    18. self.run_iter(data_batch, train_mode=True, **kwargs)
    19. ######3.
    20. self.call_hook('after_train_iter')
    21. del self.data_batch
    22. self._iter += 1
    23. ######4.
    24. self.call_hook('after_train_epoch')
    25. self._epoch += 1

     mmcv/runner/base_runner.py

    1. class BaseRunner(metaclass=ABCMeta):
    2. def call_hook(self, fn_name: str) -> None:
    3. """Call all hooks.
    4. Args:
    5. fn_name (str): The function name in each hook to be called, such as
    6. "before_train_epoch".
    7. """
    8. for hook in self._hooks:
    9. getattr(hook, fn_name)(self)

    观察代码我们可以发现,在训练的整个生命周期,有四个时间可以引入hooks,分别是before_train_epoch, before_train_iter, after_train_iter, after_train_epoch. 为什么这么命名呢?

    2.调用机制
    在训练时,利用self.call_hook执行hooks具体的操作,以OptimizerHook为例。观察call_hook函数我们发现,利用for循环调用getattr,getattr的具体作用是通过fn_name来获得属性值或getattr(hook, name)或调用同名函数getattr(hook, name)(),这里明显是后者的作用。

    现在我们解释为什么这么命名,我们可以发现,不同的hooks类在定义的时候,其主函数体是根据上面的方式唯一命名的,例如optimizer.py中的after_train_iter函数,二者一一对应,也就是说,通过这种命名结合getattr操作可以实现对hooks操作的执行。

    总的来说,这里先通过register机制将所有的hooks操作都加入self._hooks中,然后通过call_hooks中的getattr函数对self._hooks的hooks进行调用,通过命名来区分不同阶段该调用的hooks

    熟悉getattr的同学可能会有疑问,既然每种hook都有唯一的成员函数与之对应,那么我循环遍历的时候,势必会出现当前函数在某一hook类没有定义的情况,例如,在执行self.call_hook('before_train_epoch') 的时候,OptimizerHook中没有before_train_epoch函数,那getattr不是会报错吗?
    这个问题是个好问题,接下来解释原因,因为所有xxxHook都有一个父类Hook,在父类中定义了所有可能出现的方法,在子类中只需要重构需使用的函数即可,因此不会出现提到的问题,函数是存在的,只不过不执行具体操作而已。

    mmcv/runner/hooks/optimizer.py

    1. @HOOKS.register_module()
    2. class OptimizerHook(Hook):
    3. """A hook contains custom operations for the optimizer.
    4. Args:
    5. grad_clip (dict, optional): A config dict to control the clip_grad.
    6. Default: None.
    7. detect_anomalous_params (bool): This option is only used for
    8. debugging which will slow down the training speed.
    9. Detect anomalous parameters that are not included in
    10. the computational graph with `loss` as the root.
    11. There are two cases
    12. - Parameters were not used during
    13. forward pass.
    14. - Parameters were not used to produce
    15. loss.
    16. Default: False.
    17. """
    18. def __init__(self,
    19. grad_clip: Optional[dict] = None,
    20. detect_anomalous_params: bool = False):
    21. self.grad_clip = grad_clip
    22. self.detect_anomalous_params = detect_anomalous_params
    23. def clip_grads(self, params):
    24. params = list(
    25. filter(lambda p: p.requires_grad and p.grad is not None, params))
    26. if len(params) > 0:
    27. return clip_grad.clip_grad_norm_(params, **self.grad_clip)
    28. def after_train_iter(self, runner):
    29. runner.optimizer.zero_grad()
    30. if self.detect_anomalous_params:
    31. self.detect_anomalous_parameters(runner.outputs['loss'], runner)
    32. runner.outputs['loss'].backward()
    33. if self.grad_clip is not None:
    34. grad_norm = self.clip_grads(runner.model.parameters())
    35. if grad_norm is not None:
    36. # Add grad norm to the logger
    37. runner.log_buffer.update({'grad_norm': float(grad_norm)},
    38. runner.outputs['num_samples'])
    39. runner.optimizer.step()
    40. def detect_anomalous_parameters(self, loss: Tensor, runner) -> None:
    41. logger = runner.logger
    42. parameters_in_graph = set()
    43. visited = set()
    44. def traverse(grad_fn):
    45. if grad_fn is None:
    46. return
    47. if grad_fn not in visited:
    48. visited.add(grad_fn)
    49. if hasattr(grad_fn, 'variable'):
    50. parameters_in_graph.add(grad_fn.variable)
    51. parents = grad_fn.next_functions
    52. if parents is not None:
    53. for parent in parents:
    54. grad_fn = parent[0]
    55. traverse(grad_fn)
    56. traverse(loss.grad_fn)
    57. for n, p in runner.model.named_parameters():
    58. if p not in parameters_in_graph and p.requires_grad:
    59. logger.log(
    60. level=logging.ERROR,
    61. msg=f'{n} with shape {p.size()} is not '
    62. f'in the computational graph \n')


     

  • 相关阅读:
    idea2021.2.3安装炫酷插件activate-power-mode失败解决方案
    【哈佛公开课】积极心理学笔记-06乐观主义(上)
    数字孪生实际应用:智慧城市项目建设解决方案
    linux篇---解决 Linux 系统,出现“不在sudoers文件中,此事将被报告”的问题
    Linux管道与重定向
    Spring-Spring之事务底层源码解析
    LeetCode:66.加一
    隆云通空气温湿,光照三合一传感器
    12.optimizer优化器和evaluate评估
    .Net 在容器中操作宿主机
  • 原文地址:https://blog.csdn.net/qq_41368074/article/details/127111307