Skip to content

Commit 2b54e3a

Browse files
authored
feat: print optimizer param groups when verbose (#1149)
Prints which module parameters match each optimizer param group when verbose is set, both at init and on set_params. Fixes #291.
1 parent 49bae1b commit 2b54e3a

4 files changed

Lines changed: 148 additions & 2 deletions

File tree

skorch/net.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,7 @@
4141
from skorch.exceptions import SkorchAttributeError
4242
from skorch.exceptions import SkorchTrainingImpossibleError
4343
from skorch.history import History
44+
from skorch.setter import format_param_group_msg
4445
from skorch.setter import optimizer_setter
4546
from skorch.utils import _TorchLoadUnpickler
4647
from skorch.utils import _identity
@@ -2030,7 +2031,13 @@ def _get_params_for_optimizer(self, prefix, named_parameters):
20302031
matches = [i for i, (name, _) in enumerate(params) if
20312032
fnmatch.fnmatch(name, pattern)]
20322033
if matches:
2033-
p = [params.pop(i)[1] for i in reversed(matches)]
2034+
# pop high indices first so earlier indices stay valid
2035+
matched = [params.pop(i) for i in reversed(matches)]
2036+
p = [param for _, param in matched]
2037+
if self.verbose:
2038+
# show names in the order they were matched
2039+
matched_names = [name for name, _ in reversed(matched)]
2040+
print(format_param_group_msg(group, matched_names))
20342041
pgroups.append({'params': p, **group})
20352042

20362043
if params:

skorch/setter.py

Lines changed: 47 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,41 @@
11
"""Setter functions for virtual params such as ``optimizer__lr``."""
22
import re
33

4+
# param names can be arbitrarily long, keep the verbose message bounded
5+
MAX_PARAM_GROUP_MSG_LEN = 200
6+
7+
8+
def format_param_group_msg(group_config, param_names):
9+
"""Message for which module params a param group config applies to."""
10+
if not param_names:
11+
msg = (
12+
"Setting param group {} for parameters that are not among the "
13+
"module's learnable parameters (this may be unintended).".format(
14+
group_config,
15+
)
16+
)
17+
else:
18+
msg = "Setting param group {} for {}.".format(
19+
group_config,
20+
', '.join(param_names),
21+
)
22+
if len(msg) > MAX_PARAM_GROUP_MSG_LEN:
23+
msg = msg[:MAX_PARAM_GROUP_MSG_LEN - 3] + '...'
24+
return msg
25+
26+
27+
def _param_names_for_tensors(net, tensors):
28+
"""Map optimizer tensors back to module parameter names when possible."""
29+
tensor_ids = {id(t) for t in tensors}
30+
names = []
31+
get_params = getattr(net, 'get_all_learnable_params', None)
32+
if get_params is None:
33+
return names
34+
for name, p in get_params():
35+
if id(p) in tensor_ids:
36+
names.append(name)
37+
return names
38+
439

540
def _extract_optimizer_param_name_and_group(optimizer_name, param):
641
"""Extract param group and param name from the given parameter name.
@@ -44,6 +79,8 @@ def _set_optimizer_param(optimizer, param_group, param_name, value):
4479
for group in groups:
4580
group[param_name] = value
4681

82+
return groups
83+
4784

4885
def optimizer_setter(
4986
net, param, value, optimizer_attr='optimizer_', optimizer_name='optimizer'
@@ -62,9 +99,18 @@ def optimizer_setter(
6299
param_group, param_name = _extract_optimizer_param_name_and_group(
63100
optimizer_name, param)
64101

65-
_set_optimizer_param(
102+
groups = _set_optimizer_param(
66103
optimizer=getattr(net, optimizer_attr),
67104
param_group=param_group,
68105
param_name=param_name,
69106
value=value
70107
)
108+
109+
# only report for a specific param group; a global set (e.g. optimizer__lr)
110+
# touches every param and is not what #291 asks to surface
111+
if getattr(net, 'verbose', 0) and param_group != 'all':
112+
tensors = []
113+
for group in groups:
114+
tensors.extend(group.get('params', []))
115+
param_names = _param_names_for_tensors(net, tensors)
116+
print(format_param_group_msg({param_name: value}, param_names))

skorch/tests/test_net.py

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1396,6 +1396,62 @@ def test_optimizer_param_groups(self, net_cls, module_cls):
13961396
assert net.optimizer_.param_groups[1]['lr'] == 0.5
13971397
assert net.optimizer_.param_groups[2]['lr'] == net.lr
13981398

1399+
@pytest.mark.parametrize('param_groups, expected_count, expected_msgs', [
1400+
([], 0, []),
1401+
(
1402+
[('sequential.0.*', {'lr': 0.1})],
1403+
1,
1404+
[
1405+
"Setting param group {'lr': 0.1} for",
1406+
'sequential.0.weight',
1407+
'sequential.0.bias',
1408+
],
1409+
),
1410+
(
1411+
[
1412+
('sequential.0.*', {'lr': 0.1}),
1413+
('sequential.3.*', {'lr': 0.5}),
1414+
],
1415+
2,
1416+
[
1417+
"Setting param group {'lr': 0.1} for",
1418+
"Setting param group {'lr': 0.5} for",
1419+
'sequential.0.weight',
1420+
'sequential.3.weight',
1421+
],
1422+
),
1423+
])
1424+
def test_optimizer_param_groups_verbose_prints(
1425+
self, net_cls, module_cls, data, param_groups, expected_count,
1426+
expected_msgs, capsys):
1427+
# fit instead of initialize so accidental repetitions during
1428+
# training would be caught by the count
1429+
X, y = data
1430+
net = net_cls(
1431+
module_cls,
1432+
verbose=1,
1433+
max_epochs=2,
1434+
optimizer__param_groups=param_groups,
1435+
)
1436+
net.fit(X, y)
1437+
out = capsys.readouterr().out
1438+
assert out.count('Setting param group') == expected_count
1439+
for expected in expected_msgs:
1440+
assert expected in out
1441+
1442+
def test_optimizer_param_groups_silent_when_verbose_0(
1443+
self, net_cls, module_cls, capsys):
1444+
net = net_cls(
1445+
module_cls,
1446+
verbose=0,
1447+
optimizer__param_groups=[
1448+
('sequential.0.*', {'lr': 0.1}),
1449+
],
1450+
)
1451+
net.initialize()
1452+
out = capsys.readouterr().out
1453+
assert 'Setting param group' not in out
1454+
13991455
def test_module_params_in_init(self, net_cls, module_cls, data):
14001456
X, y = data
14011457

skorch/tests/test_setter.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,3 +73,40 @@ def test_only_specific_param_group_updated(self, setter, net_optim_dummy,
7373
assert updated_group_new[0][sub_param] == value
7474
assert all(old[sub_param] == new[sub_param] for old, new in zip(
7575
static_groups_pre, static_groups_new))
76+
77+
def test_set_params_verbose_prints_param_group(self, setter, capsys):
78+
from skorch import NeuralNetClassifier
79+
from skorch.toy import make_classifier
80+
81+
net = NeuralNetClassifier(make_classifier(), verbose=1, max_epochs=1)
82+
net.initialize()
83+
setter(net, 'optimizer__param_groups__0__lr', 0.03)
84+
out = capsys.readouterr().out
85+
assert "Setting param group {'lr': 0.03} for" in out
86+
assert 'sequential.0.weight' in out
87+
88+
def test_set_params_verbose_silent_for_global_param(self, setter, capsys):
89+
from skorch import NeuralNetClassifier
90+
from skorch.toy import make_classifier
91+
92+
net = NeuralNetClassifier(make_classifier(), verbose=1, max_epochs=1)
93+
net.initialize()
94+
setter(net, 'optimizer__lr', 0.03)
95+
out = capsys.readouterr().out
96+
assert 'Setting param group' not in out
97+
98+
def test_format_param_group_msg_no_known_params(self):
99+
from skorch.setter import format_param_group_msg
100+
101+
msg = format_param_group_msg({'lr': 0.1}, [])
102+
assert "{'lr': 0.1}" in msg
103+
assert 'this may be unintended' in msg
104+
105+
def test_format_param_group_msg_truncated(self):
106+
from skorch.setter import MAX_PARAM_GROUP_MSG_LEN
107+
from skorch.setter import format_param_group_msg
108+
109+
names = ['sequential.{}.weight'.format(i) for i in range(100)]
110+
msg = format_param_group_msg({'lr': 0.1}, names)
111+
assert len(msg) == MAX_PARAM_GROUP_MSG_LEN
112+
assert msg.endswith('...')

0 commit comments

Comments
 (0)