Skip to content

Commit b2c8f4d

Browse files
committed
FIX Preserve defaultdicts in to_device
1 parent 2b54e3a commit b2c8f4d

3 files changed

Lines changed: 20 additions & 1 deletion

File tree

CHANGES.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
1616

1717
### Fixed
1818

19+
- Fix moving `defaultdict` inputs to a device.
20+
1921
## [1.4.0]
2022

2123
### Added

skorch/tests/test_utils.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
"""Test for utils.py"""
22

3+
from collections import defaultdict
34
from copy import deepcopy
45
from unittest.mock import patch
56

@@ -374,6 +375,18 @@ def test_check_device_dict_torch_tensor(
374375
for k in x_dict:
375376
assert np.allclose(x_dict[k], original_x_dict[k])
376377

378+
def test_check_device_defaultdict_torch_tensor(self, to_device, x_dict):
379+
x_defaultdict = defaultdict(list, x_dict)
380+
381+
result = to_device(x_defaultdict, device='cpu')
382+
383+
assert isinstance(result, defaultdict)
384+
assert result.default_factory is list
385+
assert result['missing'] == []
386+
for key, value in x_dict.items():
387+
assert torch.equal(result[key], value)
388+
assert result[key].device.type == 'cpu'
389+
377390
@pytest.mark.parametrize('device_from, device_to', [
378391
('cpu', 'cpu'),
379392
('cpu', 'cuda'),

skorch/utils.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
55
"""
66

7+
from collections import defaultdict
78
from collections.abc import Mapping, Sequence
89
from contextlib import contextmanager
910
from enum import Enum
@@ -205,7 +206,10 @@ def to_device(X, device):
205206

206207
if isinstance(X, Mapping):
207208
# dict-like but not a dict
208-
return type(X)({key: to_device(val, device) for key, val in X.items()})
209+
mapped = {key: to_device(val, device) for key, val in X.items()}
210+
if isinstance(X, defaultdict):
211+
return type(X)(X.default_factory, mapped)
212+
return type(X)(mapped)
209213

210214
# PackedSequence class inherits from a namedtuple
211215
if isinstance(X, (tuple, list)) and (type(X) != PackedSequence):

0 commit comments

Comments
 (0)