33import re
44from collections import defaultdict
55from itertools import chain
6- from typing import Callable , Union , Dict
6+ from typing import Any , Callable , Dict , Iterator , Tuple , Type , Union
77
88import torch
99from torch import nn as nn
1313 'group_with_matcher' , 'group_modules' , 'group_parameters' , 'flatten_modules' , 'checkpoint_seq' ]
1414
1515
16- def model_parameters (model , exclude_head = False ):
16+ def model_parameters (model : nn . Module , exclude_head : bool = False ):
1717 if exclude_head :
1818 # FIXME this a bit of a quick and dirty hack to skip classifier head params based on ordering
1919 return [p for p in model .parameters ()][:- 2 ]
2020 else :
2121 return model .parameters ()
2222
2323
24- def named_apply (fn : Callable , module : nn .Module , name = '' , depth_first = True , include_root = False ) -> nn .Module :
24+ def named_apply (
25+ fn : Callable ,
26+ module : nn .Module , name = '' ,
27+ depth_first : bool = True ,
28+ include_root : bool = False ,
29+ ) -> nn .Module :
2530 if not depth_first and include_root :
2631 fn (module = module , name = name )
2732 for child_name , child_module in module .named_children ():
@@ -32,7 +37,12 @@ def named_apply(fn: Callable, module: nn.Module, name='', depth_first=True, incl
3237 return module
3338
3439
35- def named_modules (module : nn .Module , name = '' , depth_first = True , include_root = False ):
40+ def named_modules (
41+ module : nn .Module ,
42+ name : str = '' ,
43+ depth_first : bool = True ,
44+ include_root : bool = False ,
45+ ):
3646 if not depth_first and include_root :
3747 yield name , module
3848 for child_name , child_module in module .named_children ():
@@ -43,7 +53,12 @@ def named_modules(module: nn.Module, name='', depth_first=True, include_root=Fal
4353 yield name , module
4454
4555
46- def named_modules_with_params (module : nn .Module , name = '' , depth_first = True , include_root = False ):
56+ def named_modules_with_params (
57+ module : nn .Module ,
58+ name : str = '' ,
59+ depth_first : bool = True ,
60+ include_root : bool = False ,
61+ ):
4762 if module ._parameters and not depth_first and include_root :
4863 yield name , module
4964 for child_name , child_module in module .named_children ():
@@ -58,9 +73,9 @@ def named_modules_with_params(module: nn.Module, name='', depth_first=True, incl
5873
5974
6075def group_with_matcher (
61- named_objects ,
76+ named_objects : Iterator [ Tuple [ str , Any ]] ,
6277 group_matcher : Union [Dict , Callable ],
63- output_values : bool = False ,
78+ return_values : bool = False ,
6479 reverse : bool = False
6580):
6681 if isinstance (group_matcher , dict ):
@@ -96,7 +111,7 @@ def _get_grouping(name):
96111 # map layers into groups via ordinals (ints or tuples of ints) from matcher
97112 grouping = defaultdict (list )
98113 for k , v in named_objects :
99- grouping [_get_grouping (k )].append (v if output_values else k )
114+ grouping [_get_grouping (k )].append (v if return_values else k )
100115
101116 # remap to integers
102117 layer_id_to_param = defaultdict (list )
@@ -107,7 +122,7 @@ def _get_grouping(name):
107122 layer_id_to_param [lid ].extend (grouping [k ])
108123
109124 if reverse :
110- assert not output_values , "reverse mapping only sensible for name output"
125+ assert not return_values , "reverse mapping only sensible for name output"
111126 # output reverse mapping
112127 param_to_layer_id = {}
113128 for lid , lm in layer_id_to_param .items ():
@@ -121,24 +136,29 @@ def _get_grouping(name):
121136def group_parameters (
122137 module : nn .Module ,
123138 group_matcher ,
124- output_values = False ,
125- reverse = False ,
139+ return_values : bool = False ,
140+ reverse : bool = False ,
126141):
127142 return group_with_matcher (
128- module .named_parameters (), group_matcher , output_values = output_values , reverse = reverse )
143+ module .named_parameters (), group_matcher , return_values = return_values , reverse = reverse )
129144
130145
131146def group_modules (
132147 module : nn .Module ,
133148 group_matcher ,
134- output_values = False ,
135- reverse = False ,
149+ return_values : bool = False ,
150+ reverse : bool = False ,
136151):
137152 return group_with_matcher (
138- named_modules_with_params (module ), group_matcher , output_values = output_values , reverse = reverse )
153+ named_modules_with_params (module ), group_matcher , return_values = return_values , reverse = reverse )
139154
140155
141- def flatten_modules (named_modules , depth = 1 , prefix = '' , module_types = 'sequential' ):
156+ def flatten_modules (
157+ named_modules : Iterator [Tuple [str , nn .Module ]],
158+ depth : int = 1 ,
159+ prefix : Union [str , Tuple [str , ...]] = '' ,
160+ module_types : Union [str , Tuple [Type [nn .Module ]]] = 'sequential' ,
161+ ):
142162 prefix_is_tuple = isinstance (prefix , tuple )
143163 if isinstance (module_types , str ):
144164 if module_types == 'container' :
0 commit comments