Skip to content

Commit 60df5d9

Browse files
committed
distinguish between target and key in Alias
1 parent a5a08fb commit 60df5d9

1 file changed

Lines changed: 26 additions & 21 deletions

File tree

dask/_task_spec.py

Lines changed: 26 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -325,26 +325,24 @@ def resolve_aliases(dsk: dict, keys: set, dependents: dict) -> dict:
325325
raise ValueError("No keys provided")
326326
dsk = dict(dsk)
327327
work = list(keys)
328-
seen = set()
329328
while work:
330329
k = work.pop()
331-
if k in seen or k not in dsk:
330+
if k not in dsk:
332331
continue
333-
seen.add(k)
334332
t = dsk[k]
335-
if (
336-
isinstance(t, Alias)
337-
and t.key not in keys
338-
and t.key != k
339-
and t.key in dsk
340-
and len(dependents[t.key]) == 1
341-
):
342-
t = dsk[k] = dsk.pop(t.key).copy()
343-
seen.discard(k)
344-
if isinstance(t, Alias):
345-
work.append(k)
346-
else:
333+
if isinstance(t, Alias):
334+
target_key = t.target.key
335+
336+
if (
337+
target_key not in keys
338+
and target_key != k
339+
and target_key in dsk
340+
and len(dependents[target_key]) == 1
341+
):
342+
t = dsk[k] = dsk.pop(target_key).copy()
347343
t.key = k
344+
if isinstance(t, Alias):
345+
work.append(k)
348346

349347
work.extend(t.dependencies)
350348
return dsk
@@ -417,19 +415,26 @@ def __sizeof__(self) -> int:
417415

418416
class Alias(BaseTask):
419417
__weakref__: Any = None
418+
target: KeyRef
420419
__slots__ = tuple(__annotations__)
421420

422-
def __init__(self, key: KeyRef | KeyType):
423-
if isinstance(key, KeyRef):
424-
key = key.key
421+
def __init__(self, key: KeyType, target: KeyRef | KeyType | None = None):
425422
self.key = key
423+
if target is None:
424+
target = key
425+
if not isinstance(target, KeyRef):
426+
target = KeyRef(target)
427+
self.target = target
426428
self.dependencies = {key}
427429

430+
def ref(self):
431+
return self.target
432+
428433
def copy(self):
429-
return Alias(self.key)
434+
return Alias(self.key, self.target)
430435

431436
def __reduce__(self) -> str | tuple[Any, ...]:
432-
return Alias, (self.key,)
437+
return Alias, (self.key, self.target)
433438

434439
def __call__(self, values=()):
435440
self._verify_values(values)
@@ -448,7 +453,7 @@ def inline(self, dsk) -> BaseTask:
448453
return self
449454

450455
def __repr__(self):
451-
return f"Alias({self.key})"
456+
return f"Alias(key={self.key}, target={self.target})"
452457

453458
def __eq__(self, value: object) -> bool:
454459
if not isinstance(value, Alias):

0 commit comments

Comments
 (0)