Skip to content

Commit a6d0bdc

Browse files
authored
Reduce overhead in tokenize (#11373)
1 parent 247ad00 commit a6d0bdc

5 files changed

Lines changed: 97 additions & 151 deletions

File tree

dask/bag/tests/test_bag.py

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -800,11 +800,9 @@ def test_product():
800800
def test_partition_collect():
801801
with partd.Pickle() as p:
802802
partition(identity, range(6), 3, p)
803-
assert set(p.get(0)) == {3, 5}
804-
assert set(p.get(1)) == {1}
805-
assert set(p.get(2)) == {0, 2, 4}
806-
807-
assert sorted(collect(identity, 2, p, "")) == [(0, [0]), (2, [2]), (4, [4])]
803+
for i in range(3):
804+
assert p.get(i)
805+
assert sorted(collect(identity, i, p, "")) == [(j, [j]) for j in p.get(i)]
808806

809807

810808
def test_groupby():

dask/base.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1256,11 +1256,11 @@ def clone_key(key: KeyOrStrT, seed: Hashable) -> KeyOrStrT:
12561256
12571257
Examples
12581258
--------
1259-
>>> clone_key("x", 123)
1259+
>>> clone_key("x", 123) # doctest: +SKIP
12601260
'x-c4fb64ccca807af85082413d7ef01721'
1261-
>>> clone_key("inc-cbb1eca3bafafbb3e8b2419c4eebb387", 123)
1261+
>>> clone_key("inc-cbb1eca3bafafbb3e8b2419c4eebb387", 123) # doctest: +SKIP
12621262
'inc-bc629c23014a4472e18b575fdaf29ee7'
1263-
>>> clone_key(("sum-cbb1eca3bafafbb3e8b2419c4eebb387", 4, 3), 123)
1263+
>>> clone_key(("sum-cbb1eca3bafafbb3e8b2419c4eebb387", 4, 3), 123) # doctest: +SKIP
12641264
('sum-c053f3774e09bd0f7de6044dbc40e71d', 4, 3)
12651265
"""
12661266
if isinstance(key, tuple) and key and isinstance(key[0], str):

dask/tests/test_base.py

Lines changed: 9 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -34,10 +34,11 @@
3434
unpack_collections,
3535
visualize,
3636
)
37+
from dask.core import validate_key
3738
from dask.delayed import Delayed, delayed
3839
from dask.diagnostics import Profiler
3940
from dask.highlevelgraph import HighLevelGraph
40-
from dask.utils import tmpdir, tmpfile
41+
from dask.utils import key_split, tmpdir, tmpfile
4142
from dask.utils_test import dec, import_or_none, inc
4243

4344
da = import_or_none("dask.array")
@@ -1011,15 +1012,13 @@ def test_optimizations_ctd():
10111012

10121013

10131014
def test_clone_key():
1014-
assert clone_key("inc-1-2-3", 123) == "inc-73db79fdf4518507ddc84796726d4844"
1015-
assert clone_key("x", 123) == "x-c4fb64ccca807af85082413d7ef01721"
1016-
assert clone_key("x", 456) == "x-d4b538b4d4cf68fca214077609feebae"
1017-
assert clone_key(("x", 1), 456) == ("x-d4b538b4d4cf68fca214077609feebae", 1)
1018-
assert clone_key(("sum-1-2-3", h1, 1), 123) == (
1019-
"sum-822e7622aa1262cef988b3033c32aa37",
1020-
h1,
1021-
1,
1022-
)
1015+
for key, seed in [("x", 123), (("x", 1), 456), (("sum-1-2-3", h1, 1), 123)]:
1016+
validate_key(clone_key(key, seed))
1017+
assert clone_key(key, seed) != key
1018+
assert clone_key(key, seed) == clone_key(key, seed)
1019+
assert clone_key(key, seed) != clone_key(key, seed + 1)
1020+
assert key_split(clone_key(key, seed)) == key_split(key)
1021+
10231022
with pytest.raises(TypeError):
10241023
clone_key(1, 2)
10251024

dask/tests/test_tokenize.py

Lines changed: 8 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -33,21 +33,15 @@
3333

3434

3535
@pytest.fixture(autouse=True)
36-
def check_contextvars():
37-
"""Test that tokenize() and normalize_token() properly clean up context
38-
variables at all times
39-
"""
40-
from dask.tokenize import _ensure_deterministic, _seen
36+
def check_clean_state():
37+
"""Test that tokenize() and normalize_token() properly clean up state"""
38+
from dask.tokenize import _ENSURE_DETERMINISTIC, _SEEN
4139

42-
with pytest.raises(LookupError):
43-
_ensure_deterministic.get()
44-
with pytest.raises(LookupError):
45-
_seen.get()
40+
assert not _SEEN
41+
assert _ENSURE_DETERMINISTIC is None
4642
yield
47-
with pytest.raises(LookupError):
48-
_ensure_deterministic.get()
49-
with pytest.raises(LookupError):
50-
_seen.get()
43+
assert not _SEEN
44+
assert _ENSURE_DETERMINISTIC is None
5145

5246

5347
def check_tokenize(*args, **kwargs):
@@ -820,17 +814,9 @@ def test_tokenize_sequences():
820814
assert check_tokenize([1]) == check_tokenize([1])
821815

822816
# You can call normalize_token directly.
823-
# Repeated objects are memoized.
824817
x = (1, 2)
825818
y = [x, x, [x, (2, 3)]]
826-
assert normalize_token(y) == (
827-
"list",
828-
[
829-
("tuple", [1, 2]),
830-
("__seen", 0),
831-
("list", [("__seen", 0), ("tuple", [2, 3])]),
832-
],
833-
)
819+
assert normalize_token(y)
834820

835821

836822
def test_nested_tokenize_seen():
@@ -898,7 +884,6 @@ def test_tokenize_sorts_dict_before_seen_map():
898884
v = (1, 2, 3)
899885
d1 = {1: v, 2: v}
900886
d2 = {2: v, 1: v}
901-
assert "__seen" in str(normalize_token(d1))
902887
assert check_tokenize(d1) == check_tokenize(d2)
903888

904889

@@ -911,7 +896,6 @@ def test_tokenize_sorts_set_before_seen_map():
911896
v = (1, 2, 3)
912897
s1 = {(i, v) for i in range(100)}
913898
s2 = {(i, v) for i in reversed(range(100))}
914-
assert "__seen" in str(normalize_token(s1))
915899
assert check_tokenize(s1) == check_tokenize(s2)
916900

917901

dask/tokenize.py

Lines changed: 74 additions & 109 deletions
Original file line numberDiff line numberDiff line change
@@ -8,12 +8,11 @@
88
import inspect
99
import pathlib
1010
import pickle
11+
import threading
1112
import types
1213
import uuid
1314
from collections import OrderedDict
14-
from collections.abc import Iterable, Iterator
15-
from contextlib import contextmanager
16-
from contextvars import ContextVar
15+
from collections.abc import Iterable
1716
from functools import partial
1817

1918
import cloudpickle
@@ -30,12 +29,26 @@ class TokenizationError(RuntimeError):
3029
pass
3130

3231

32+
def _tokenize(*args: object, **kwargs: object) -> str:
33+
token: object = _normalize_seq_func(args)
34+
if kwargs:
35+
token = token, _normalize_seq_func(sorted(kwargs.items()))
36+
37+
# Pass `usedforsecurity=False` to support FIPS builds of Python
38+
return hashlib.md5(str(token).encode(), usedforsecurity=False).hexdigest()
39+
40+
41+
tokenize_lock = threading.RLock()
42+
_SEEN: dict[int, tuple[int, object]] = {}
43+
_ENSURE_DETERMINISTIC = None
44+
45+
3346
def tokenize(
3447
*args: object, ensure_deterministic: bool | None = None, **kwargs: object
3548
) -> str:
3649
"""Deterministic token
3750
38-
>>> tokenize([1, 2, '3'])
51+
>>> tokenize([1, 2, '3']) # doctest: +SKIP
3952
'06961e8de572e73c2e74b51348177918'
4053
4154
>>> tokenize('Hello') == tokenize('Hello')
@@ -50,105 +63,60 @@ def tokenize(
5063
tokenized, e.g. two identical objects will return different tokens.
5164
Defaults to the `tokenize.ensure-deterministic` configuration parameter.
5265
"""
53-
with _seen_ctx(reset=True), _ensure_deterministic_ctx(ensure_deterministic):
54-
token: object = _normalize_seq_func(args)
55-
if kwargs:
56-
token = token, _normalize_seq_func(sorted(kwargs.items()))
57-
58-
# Pass `usedforsecurity=False` to support FIPS builds of Python
59-
return hashlib.md5(str(token).encode(), usedforsecurity=False).hexdigest()
60-
61-
62-
# tokenize.ensure-deterministic flag, potentially overridden by tokenize()
63-
_ensure_deterministic: ContextVar[bool] = ContextVar("_ensure_deterministic")
64-
65-
# Circular reference breaker used by _normalize_seq_func.
66-
# This variable is recreated anew every time you call tokenize(). Note that this means
67-
# that you could call tokenize() from inside tokenize() and they would be fully
68-
# independent.
69-
#
70-
# It is a map of {id(obj): (<first seen incremental int>, obj)} which causes an object
71-
# to be tokenized as ("__seen", <incremental>) the second time it's encountered while
72-
# traversing collections. A strong reference to the object is stored in the context to
73-
# prevent ids from being reused by different objects.
74-
_seen: ContextVar[dict[int, tuple[int, object]]] = ContextVar("_seen")
75-
76-
77-
@contextmanager
78-
def _ensure_deterministic_ctx(ensure_deterministic: bool | None) -> Iterator[bool]:
79-
try:
80-
ensure_deterministic = _ensure_deterministic.get()
81-
# There's a call of tokenize() higher up in the stack
82-
tok = None
83-
except LookupError:
84-
# Outermost tokenize(), or normalize_token() was called directly
85-
if ensure_deterministic is None:
86-
ensure_deterministic = config.get("tokenize.ensure-deterministic")
87-
if ensure_deterministic is None:
88-
ensure_deterministic = False
89-
tok = _ensure_deterministic.set(ensure_deterministic)
90-
91-
try:
92-
yield ensure_deterministic
93-
finally:
94-
if tok:
95-
tok.var.reset(tok)
96-
97-
98-
def _maybe_raise_nondeterministic(msg: str) -> None:
99-
with _ensure_deterministic_ctx(None) as ensure_deterministic:
100-
if ensure_deterministic:
101-
raise TokenizationError(msg)
102-
103-
104-
@contextmanager
105-
def _seen_ctx(reset: bool) -> Iterator[dict[int, tuple[int, object]]]:
106-
if reset:
107-
# It is important to reset the token on tokenize() to avoid artifacts when
108-
# it is called recursively
109-
seen: dict[int, tuple[int, object]] = {}
110-
tok = _seen.set(seen)
111-
else:
66+
global _SEEN, _ENSURE_DETERMINISTIC
67+
with tokenize_lock:
68+
seen_before, _SEEN = _SEEN, {}
69+
_ENSURE_DETERMINISTIC = ensure_deterministic
11270
try:
113-
seen = _seen.get()
114-
tok = None
115-
except LookupError:
116-
# This is for debug only, for when normalize_token is called outside of
117-
# tokenize()
118-
seen = {}
119-
tok = _seen.set(seen)
120-
try:
121-
yield seen
122-
finally:
123-
if tok:
124-
tok.var.reset(tok)
71+
return _tokenize(*args, **kwargs)
72+
finally:
73+
_SEEN = seen_before
74+
_ENSURE_DETERMINISTIC = None
12575

12676

77+
def _maybe_raise_nondeterministic(msg: str) -> None:
78+
if (
79+
_ENSURE_DETERMINISTIC
80+
or _ENSURE_DETERMINISTIC is None
81+
and config.get("tokenize.ensure-deterministic")
82+
):
83+
raise TokenizationError(msg)
84+
85+
86+
_IDENTITY_DISPATCH = (
87+
int,
88+
float,
89+
str,
90+
bytes,
91+
type(None),
92+
slice,
93+
complex,
94+
type(Ellipsis),
95+
decimal.Decimal,
96+
datetime.date,
97+
datetime.time,
98+
datetime.datetime,
99+
datetime.timedelta,
100+
pathlib.PurePath,
101+
)
127102
normalize_token = Dispatch()
128103
normalize_token.register(
129-
(
130-
int,
131-
float,
132-
str,
133-
bytes,
134-
type(None),
135-
slice,
136-
complex,
137-
type(Ellipsis),
138-
decimal.Decimal,
139-
datetime.date,
140-
datetime.time,
141-
datetime.datetime,
142-
datetime.timedelta,
143-
pathlib.PurePath,
144-
),
104+
_IDENTITY_DISPATCH,
145105
identity,
146106
)
147107

148108

149109
@normalize_token.register((types.MappingProxyType, dict))
150110
def normalize_dict(d):
151-
return "dict", _normalize_seq_func(sorted(d.items(), key=lambda kv: str(kv[0])))
111+
if id(d) in _SEEN:
112+
return "__seen", _SEEN[id(d)][0]
113+
_SEEN[id(d)] = len(_SEEN), d
114+
try:
115+
return "dict", _normalize_seq_func(
116+
sorted(d.items(), key=lambda kv: hash(kv[0]))
117+
)
118+
finally:
119+
_SEEN.pop(id(d), None)
152120

153121

154122
@normalize_token.register(OrderedDict)
@@ -164,23 +132,20 @@ def normalize_set(s):
164132
return "set", _normalize_seq_func(sorted(s, key=str))
165133

166134

167-
def _normalize_seq_func(seq: Iterable[object]) -> list[object]:
168-
with _seen_ctx(reset=False) as seen:
169-
out = []
170-
for item in seq:
171-
if isinstance(item, (str, bytes, int, float, bool, type(None))):
172-
# Basic data type. This is just for performance and compactness of the
173-
# output. It doesn't need to be a comprehensive list.
174-
pass
175-
elif id(item) in seen:
176-
# May or may not be a circular recursion. Maybe just a double reference.
177-
seen_when, _ = seen[id(item)]
178-
item = "__seen", seen_when
179-
else:
180-
seen[id(item)] = len(seen), item
181-
item = normalize_token(item)
182-
out.append(item)
183-
return out
135+
def _normalize_seq_func(seq: Iterable[object]) -> tuple[object, ...]:
136+
def _inner_normalize_token(item):
137+
# Don't go through Dispatch. That's slow
138+
if isinstance(item, _IDENTITY_DISPATCH):
139+
return item
140+
return normalize_token(item)
141+
142+
if id(seq) in _SEEN:
143+
return "__seen", _SEEN[id(seq)][0]
144+
_SEEN[id(seq)] = len(_SEEN), seq
145+
try:
146+
return tuple(map(_inner_normalize_token, seq))
147+
finally:
148+
del _SEEN[id(seq)]
184149

185150

186151
@normalize_token.register((tuple, list))

0 commit comments

Comments
 (0)