Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/dask/_expr.py: 22%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

761 statements  

1from __future__ import annotations 

2 

3import functools 

4import os 

5import uuid 

6import warnings 

7import weakref 

8from collections import defaultdict 

9from collections.abc import Generator 

10from typing import TYPE_CHECKING, Any, Literal, TypeAlias 

11 

12import toolz 

13 

14import dask 

15from dask._task_spec import Task, convert_legacy_graph 

16from dask.tokenize import _tokenize_deterministic 

17from dask.typing import Key 

18from dask.utils import ensure_dict, funcname, import_required 

19 

20if TYPE_CHECKING: 

21 from dask.highlevelgraph import HighLevelGraph 

22 

23OptimizerStage: TypeAlias = Literal[ 

24 "logical", 

25 "simplified-logical", 

26 "tuned-logical", 

27 "physical", 

28 "simplified-physical", 

29 "fused", 

30] 

31 

32 

33def _unpack_collections(o): 

34 from dask.delayed import Delayed 

35 

36 if isinstance(o, Expr): 

37 return o 

38 

39 if not isinstance(o, Delayed): 

40 try: 

41 expr = o.expr 

42 except AttributeError: 

43 pass 

44 else: 

45 if isinstance(expr, Expr): 

46 return expr 

47 return o 

48 

49 

50class Expr: 

51 _parameters: list[str] = [] 

52 _defaults: dict[str, Any] = {} 

53 

54 _pickle_functools_cache: bool = True 

55 

56 operands: list 

57 

58 _determ_token: str | None 

59 

60 def __new__(cls, *args, _determ_token=None, **kwargs): 

61 operands = list(args) 

62 for parameter in cls._parameters[len(operands) :]: 

63 try: 

64 operands.append(kwargs.pop(parameter)) 

65 except KeyError: 

66 operands.append(cls._defaults[parameter]) 

67 assert not kwargs, kwargs 

68 inst = object.__new__(cls) 

69 

70 inst._determ_token = _determ_token 

71 inst.operands = [_unpack_collections(o) for o in operands] 

72 # This is typically cached. Make sure the cache is populated by calling 

73 # it once 

74 inst._name 

75 return inst 

76 

77 def _tune_down(self): 

78 return None 

79 

80 def _tune_up(self, parent): 

81 return None 

82 

83 def finalize_compute(self): 

84 return self 

85 

86 def _operands_for_repr(self): 

87 return [f"{param}={op!r}" for param, op in zip(self._parameters, self.operands)] 

88 

89 def __str__(self): 

90 s = ", ".join(self._operands_for_repr()) 

91 return f"{type(self).__name__}({s})" 

92 

93 def __repr__(self): 

94 return str(self) 

95 

96 def _tree_repr_argument_construction(self, i, op, header): 

97 try: 

98 param = self._parameters[i] 

99 default = self._defaults[param] 

100 except (IndexError, KeyError): 

101 param = self._parameters[i] if i < len(self._parameters) else "" 

102 default = "--no-default--" 

103 

104 if repr(op) != repr(default): 

105 if param: 

106 header += f" {param}={op!r}" 

107 else: 

108 header += repr(op) 

109 return header 

110 

111 def _tree_repr_lines(self, indent=0, recursive=True): 

112 return " " * indent + repr(self) 

113 

114 def tree_repr(self): 

115 return os.linesep.join(self._tree_repr_lines()) 

116 

117 def analyze(self, filename: str | None = None, format: str | None = None) -> None: 

118 from dask.dataframe.dask_expr._expr import Expr as DFExpr 

119 from dask.dataframe.dask_expr.diagnostics import analyze 

120 

121 if not isinstance(self, DFExpr): 

122 raise TypeError( 

123 "analyze is only supported for dask.dataframe.Expr objects." 

124 ) 

125 return analyze(self, filename=filename, format=format) 

126 

127 def explain( 

128 self, stage: OptimizerStage = "fused", format: str | None = None 

129 ) -> None: 

130 from dask.dataframe.dask_expr.diagnostics import explain 

131 

132 return explain(self, stage, format) 

133 

134 def pprint(self): 

135 for line in self._tree_repr_lines(): 

136 print(line) 

137 

138 def __hash__(self): 

139 return hash(self._name) 

140 

141 def __dask_tokenize__(self): 

142 if not self._determ_token: 

143 # If the subclass does not implement a __dask_tokenize__ we'll want 

144 # to tokenize all operands. 

145 # Note how this differs to the implementation of 

146 # Expr.deterministic_token 

147 self._determ_token = _tokenize_deterministic(type(self), *self.operands) 

148 return self._determ_token 

149 

150 def __dask_keys__(self): 

151 """The keys for this expression 

152 

153 This is used to determine the keys of the output collection 

154 when this expression is computed. 

155 

156 Returns 

157 ------- 

158 keys: list 

159 The keys for this expression 

160 """ 

161 return [(self._name, i) for i in range(self.npartitions)] 

162 

163 @staticmethod 

164 def _reconstruct(*args): 

165 typ, *operands, token, cache = args 

166 inst = typ(*operands, _determ_token=token) 

167 for k, v in cache.items(): 

168 inst.__dict__[k] = v 

169 return inst 

170 

171 def __reduce__(self): 

172 if dask.config.get("dask-expr-no-serialize", False): 

173 raise RuntimeError(f"Serializing a {type(self)} object") 

174 cache = {} 

175 if type(self)._pickle_functools_cache: 

176 for k, v in type(self).__dict__.items(): 

177 if isinstance(v, functools.cached_property) and k in self.__dict__: 

178 cache[k] = getattr(self, k) 

179 

180 return Expr._reconstruct, ( 

181 type(self), 

182 *self.operands, 

183 self.deterministic_token, 

184 cache, 

185 ) 

186 

187 def _depth(self, cache=None): 

188 """Depth of the expression tree 

189 

190 Returns 

191 ------- 

192 depth: int 

193 """ 

194 if cache is None: 

195 cache = {} 

196 if not self.dependencies(): 

197 return 1 

198 else: 

199 result = [] 

200 for expr in self.dependencies(): 

201 if expr._name in cache: 

202 result.append(cache[expr._name]) 

203 else: 

204 result.append(expr._depth(cache) + 1) 

205 cache[expr._name] = result[-1] 

206 return max(result) 

207 

208 def __setattr__(self, name: str, value: Any) -> None: 

209 if name in ["operands", "_determ_token"]: 

210 object.__setattr__(self, name, value) 

211 return 

212 try: 

213 params = type(self)._parameters 

214 operands = object.__getattribute__(self, "operands") 

215 operands[params.index(name)] = value 

216 except ValueError: 

217 raise AttributeError( 

218 f"{type(self).__name__} object has no attribute {name}" 

219 ) 

220 

221 def operand(self, key): 

222 # Access an operand unambiguously 

223 # (e.g. if the key is reserved by a method/property) 

224 return self.operands[type(self)._parameters.index(key)] 

225 

226 def dependencies(self): 

227 # Dependencies are `Expr` operands only 

228 return [operand for operand in self.operands if isinstance(operand, Expr)] 

229 

230 def _task(self, key: Key, index: int) -> Task: 

231 """The task for the i'th partition 

232 

233 Parameters 

234 ---------- 

235 index: 

236 The index of the partition of this dataframe 

237 

238 Examples 

239 -------- 

240 >>> class Add(Expr): 

241 ... def _task(self, i): 

242 ... return Task( 

243 ... self.__dask_keys__()[i], 

244 ... operator.add, 

245 ... TaskRef((self.left._name, i)), 

246 ... TaskRef((self.right._name, i)) 

247 ... ) 

248 

249 Returns 

250 ------- 

251 task: 

252 The Dask task to compute this partition 

253 

254 See Also 

255 -------- 

256 Expr._layer 

257 """ 

258 raise NotImplementedError( 

259 "Expressions should define either _layer (full dictionary) or _task" 

260 f" (single task). This expression {type(self)} defines neither" 

261 ) 

262 

263 def _layer(self) -> dict: 

264 """The graph layer added by this expression. 

265 

266 Simple expressions that apply one task per partition can choose to only 

267 implement `Expr._task` instead. 

268 

269 Examples 

270 -------- 

271 >>> class Add(Expr): 

272 ... def _layer(self): 

273 ... return { 

274 ... name: Task( 

275 ... name, 

276 ... operator.add, 

277 ... TaskRef((self.left._name, i)), 

278 ... TaskRef((self.right._name, i)) 

279 ... ) 

280 ... for i, name in enumerate(self.__dask_keys__()) 

281 ... } 

282 

283 Returns 

284 ------- 

285 layer: dict 

286 The Dask task graph added by this expression 

287 

288 See Also 

289 -------- 

290 Expr._task 

291 Expr.__dask_graph__ 

292 """ 

293 

294 return { 

295 (self._name, i): self._task((self._name, i), i) 

296 for i in range(self.npartitions) 

297 } 

298 

299 def rewrite(self, kind: str, rewritten): 

300 """Rewrite an expression 

301 

302 This leverages the ``._{kind}_down`` and ``._{kind}_up`` 

303 methods defined on each class 

304 

305 Returns 

306 ------- 

307 expr: 

308 output expression 

309 changed: 

310 whether or not any change occurred 

311 """ 

312 if self._name in rewritten: 

313 return rewritten[self._name] 

314 

315 expr = self 

316 down_name = f"_{kind}_down" 

317 up_name = f"_{kind}_up" 

318 while True: 

319 _continue = False 

320 

321 # Rewrite this node 

322 out = getattr(expr, down_name)() 

323 if out is None: 

324 out = expr 

325 if not isinstance(out, Expr): 

326 return out 

327 if out._name != expr._name: 

328 expr = out 

329 continue 

330 

331 # Allow children to rewrite their parents 

332 for child in expr.dependencies(): 

333 out = getattr(child, up_name)(expr) 

334 if out is None: 

335 out = expr 

336 if not isinstance(out, Expr): 

337 return out 

338 if out is not expr and out._name != expr._name: 

339 expr = out 

340 _continue = True 

341 break 

342 

343 if _continue: 

344 continue 

345 

346 # Rewrite all of the children 

347 new_operands = [] 

348 changed = False 

349 for operand in expr.operands: 

350 if isinstance(operand, Expr): 

351 new = operand.rewrite(kind=kind, rewritten=rewritten) 

352 rewritten[operand._name] = new 

353 if new._name != operand._name: 

354 changed = True 

355 else: 

356 new = operand 

357 new_operands.append(new) 

358 

359 if changed: 

360 expr = type(expr)(*new_operands) 

361 continue 

362 else: 

363 break 

364 

365 return expr 

366 

367 def simplify_once(self, dependents: defaultdict, simplified: dict): 

368 """Simplify an expression 

369 

370 This leverages the ``._simplify_down`` and ``._simplify_up`` 

371 methods defined on each class 

372 

373 Parameters 

374 ---------- 

375 

376 dependents: defaultdict[list] 

377 The dependents for every node. 

378 simplified: dict 

379 Cache of simplified expressions for these dependents. 

380 

381 Returns 

382 ------- 

383 expr: 

384 output expression 

385 """ 

386 # Check if we've already simplified for these dependents 

387 if self._name in simplified: 

388 return simplified[self._name] 

389 

390 expr = self 

391 

392 while True: 

393 out = expr._simplify_down() 

394 if out is None: 

395 out = expr 

396 if not isinstance(out, Expr): 

397 return out 

398 if out._name != expr._name: 

399 expr = out 

400 

401 # Allow children to simplify their parents 

402 for child in expr.dependencies(): 

403 out = child._simplify_up(expr, dependents) 

404 if out is None: 

405 out = expr 

406 

407 if not isinstance(out, Expr): 

408 return out 

409 if out is not expr and out._name != expr._name: 

410 expr = out 

411 break 

412 

413 # Rewrite all of the children 

414 new_operands = [] 

415 changed = False 

416 for operand in expr.operands: 

417 if isinstance(operand, Expr): 

418 # Bandaid for now, waiting for Singleton 

419 dependents[operand._name].append(weakref.ref(expr)) 

420 new = operand.simplify_once( 

421 dependents=dependents, simplified=simplified 

422 ) 

423 simplified[operand._name] = new 

424 if new._name != operand._name: 

425 changed = True 

426 else: 

427 new = operand 

428 new_operands.append(new) 

429 

430 if changed: 

431 expr = type(expr)(*new_operands) 

432 

433 break 

434 

435 return expr 

436 

437 def optimize(self, fuse: bool = False) -> Expr: 

438 stage: OptimizerStage = "fused" if fuse else "simplified-physical" 

439 

440 return optimize_until(self, stage) 

441 

442 def fuse(self) -> Expr: 

443 return self 

444 

445 def simplify(self) -> Expr: 

446 expr = self 

447 seen = set() 

448 while True: 

449 dependents = collect_dependents(expr) 

450 new = expr.simplify_once(dependents=dependents, simplified={}) 

451 if new._name == expr._name: 

452 break 

453 if new._name in seen: 

454 raise RuntimeError( 

455 f"Optimizer does not converge. {expr!r} simplified to {new!r} which was already seen. " 

456 "Please report this issue on the dask issue tracker with a minimal reproducer." 

457 ) 

458 seen.add(new._name) 

459 expr = new 

460 return expr 

461 

462 def _simplify_down(self): 

463 return 

464 

465 def _simplify_up(self, parent, dependents): 

466 return 

467 

468 def lower_once(self, lowered: dict): 

469 # Check for a cached result 

470 try: 

471 return lowered[self._name] 

472 except KeyError: 

473 pass 

474 

475 expr = self 

476 

477 # Lower this node 

478 out = expr._lower() 

479 if out is None: 

480 out = expr 

481 if not isinstance(out, Expr): 

482 return out 

483 

484 # Lower all children 

485 new_operands = [] 

486 changed = False 

487 for operand in out.operands: 

488 if isinstance(operand, Expr): 

489 new = operand.lower_once(lowered) 

490 if new._name != operand._name: 

491 changed = True 

492 else: 

493 new = operand 

494 new_operands.append(new) 

495 

496 if changed: 

497 out = type(out)(*new_operands) 

498 

499 # Cache the result and return 

500 return lowered.setdefault(self._name, out) 

501 

502 def lower_completely(self) -> Expr: 

503 """Lower an expression completely 

504 

505 This calls the ``lower_once`` method in a loop 

506 until nothing changes. This function does not 

507 apply any other optimizations (like ``simplify``). 

508 

509 Returns 

510 ------- 

511 expr: 

512 output expression 

513 

514 See Also 

515 -------- 

516 Expr.lower_once 

517 Expr._lower 

518 """ 

519 # Lower until nothing changes 

520 expr = self 

521 lowered: dict = {} 

522 while True: 

523 new = expr.lower_once(lowered) 

524 if new._name == expr._name: 

525 break 

526 expr = new 

527 return expr 

528 

529 def _lower(self): 

530 return 

531 

532 @functools.cached_property 

533 def _funcname(self) -> str: 

534 return funcname(type(self)).lower() 

535 

536 @property 

537 def deterministic_token(self): 

538 if not self._determ_token: 

539 # Just tokenize self to fall back on __dask_tokenize__ 

540 # Note how this differs to the implementation of __dask_tokenize__ 

541 self._determ_token = self.__dask_tokenize__() 

542 return self._determ_token 

543 

544 @functools.cached_property 

545 def _name(self) -> str: 

546 return f"{self._funcname}-{self.deterministic_token}" 

547 

548 @property 

549 def _meta(self): 

550 raise NotImplementedError() 

551 

552 @classmethod 

553 def _annotations_tombstone(cls) -> _AnnotationsTombstone: 

554 return _AnnotationsTombstone() 

555 

556 def __dask_annotations__(self): 

557 return {} 

558 

559 def __dask_graph__(self): 

560 """Traverse expression tree, collect layers 

561 

562 Subclasses generally do not want to override this method unless custom 

563 logic is required to treat (e.g. ignore) specific operands during graph 

564 generation. 

565 

566 See also 

567 -------- 

568 Expr._layer 

569 Expr._task 

570 """ 

571 stack = [self] 

572 seen = set() 

573 layers = [] 

574 while stack: 

575 expr = stack.pop() 

576 

577 if expr._name in seen: 

578 continue 

579 seen.add(expr._name) 

580 

581 layers.append(expr._layer()) 

582 for operand in expr.dependencies(): 

583 stack.append(operand) 

584 

585 return toolz.merge(layers) 

586 

587 @property 

588 def dask(self): 

589 return self.__dask_graph__() 

590 

591 def substitute(self, old, new) -> Expr: 

592 """Substitute a specific term within the expression 

593 

594 Note that replacing non-`Expr` terms may produce 

595 unexpected results, and is not recommended. 

596 Substituting boolean values is not allowed. 

597 

598 Parameters 

599 ---------- 

600 old: 

601 Old term to find and replace. 

602 new: 

603 New term to replace instances of `old` with. 

604 

605 Examples 

606 -------- 

607 >>> (df + 10).substitute(10, 20) # doctest: +SKIP 

608 df + 20 

609 """ 

610 return self._substitute(old, new, _seen=set()) 

611 

612 def _substitute(self, old, new, _seen): 

613 if self._name in _seen: 

614 return self 

615 # Check if we are replacing a literal 

616 if isinstance(old, Expr): 

617 substitute_literal = False 

618 if self._name == old._name: 

619 return new 

620 else: 

621 substitute_literal = True 

622 if isinstance(old, bool): 

623 raise TypeError("Arguments to `substitute` cannot be bool.") 

624 

625 new_exprs = [] 

626 update = False 

627 for operand in self.operands: 

628 if isinstance(operand, Expr): 

629 val = operand._substitute(old, new, _seen) 

630 if operand._name != val._name: 

631 update = True 

632 new_exprs.append(val) 

633 elif ( 

634 "Fused" in type(self).__name__ 

635 and isinstance(operand, list) 

636 and all(isinstance(op, Expr) for op in operand) 

637 ): 

638 # Special handling for `Fused`. 

639 # We make no promise to dive through a 

640 # list operand in general, but NEED to 

641 # do so for the `Fused.exprs` operand. 

642 val = [] 

643 for op in operand: 

644 val.append(op._substitute(old, new, _seen)) 

645 if val[-1]._name != op._name: 

646 update = True 

647 new_exprs.append(val) 

648 elif ( 

649 substitute_literal 

650 and not isinstance(operand, bool) 

651 and isinstance(operand, type(old)) 

652 and operand == old 

653 ): 

654 new_exprs.append(new) 

655 update = True 

656 else: 

657 new_exprs.append(operand) 

658 

659 if update: # Only recreate if something changed 

660 return type(self)(*new_exprs) 

661 else: 

662 _seen.add(self._name) 

663 return self 

664 

665 def substitute_parameters(self, substitutions: dict) -> Expr: 

666 """Substitute specific `Expr` parameters 

667 

668 Parameters 

669 ---------- 

670 substitutions: 

671 Mapping of parameter keys to new values. Keys that 

672 are not found in ``self._parameters`` will be ignored. 

673 """ 

674 if not substitutions: 

675 return self 

676 

677 changed = False 

678 new_operands = [] 

679 for i, operand in enumerate(self.operands): 

680 if i < len(self._parameters) and self._parameters[i] in substitutions: 

681 new_operands.append(substitutions[self._parameters[i]]) 

682 changed = True 

683 else: 

684 new_operands.append(operand) 

685 if changed: 

686 return type(self)(*new_operands) 

687 return self 

688 

689 def _node_label_args(self): 

690 """Operands to include in the node label by `visualize`""" 

691 return self.dependencies() 

692 

693 def _to_graphviz( 

694 self, 

695 rankdir="BT", 

696 graph_attr=None, 

697 node_attr=None, 

698 edge_attr=None, 

699 **kwargs, 

700 ): 

701 from dask.dot import label, name 

702 

703 graphviz = import_required( 

704 "graphviz", 

705 "Drawing dask graphs with the graphviz visualization engine requires the `graphviz` " 

706 "python library and the `graphviz` system library.\n\n" 

707 "Please either conda or pip install as follows:\n\n" 

708 " conda install python-graphviz # either conda install\n" 

709 " python -m pip install graphviz # or pip install and follow installation instructions", 

710 ) 

711 

712 graph_attr = graph_attr or {} 

713 node_attr = node_attr or {} 

714 edge_attr = edge_attr or {} 

715 

716 graph_attr["rankdir"] = rankdir 

717 node_attr["shape"] = "box" 

718 node_attr["fontname"] = "helvetica" 

719 

720 graph_attr.update(kwargs) 

721 g = graphviz.Digraph( 

722 graph_attr=graph_attr, 

723 node_attr=node_attr, 

724 edge_attr=edge_attr, 

725 ) 

726 

727 stack = [self] 

728 seen = set() 

729 dependencies = {} 

730 while stack: 

731 expr = stack.pop() 

732 

733 if expr._name in seen: 

734 continue 

735 seen.add(expr._name) 

736 

737 dependencies[expr] = set(expr.dependencies()) 

738 for dep in expr.dependencies(): 

739 stack.append(dep) 

740 

741 cache = {} 

742 for expr in dependencies: 

743 expr_name = name(expr) 

744 attrs = {} 

745 

746 # Make node label 

747 deps = [ 

748 funcname(type(dep)) if isinstance(dep, Expr) else str(dep) 

749 for dep in expr._node_label_args() 

750 ] 

751 _label = funcname(type(expr)) 

752 if deps: 

753 _label = f"{_label}({', '.join(deps)})" if deps else _label 

754 node_label = label(_label, cache=cache) 

755 

756 attrs.setdefault("label", str(node_label)) 

757 attrs.setdefault("fontsize", "20") 

758 g.node(expr_name, **attrs) 

759 

760 for expr, deps in dependencies.items(): 

761 expr_name = name(expr) 

762 for dep in deps: 

763 dep_name = name(dep) 

764 g.edge(dep_name, expr_name) 

765 

766 return g 

767 

768 def visualize(self, filename="dask-expr.svg", format=None, **kwargs): 

769 """ 

770 Visualize the expression graph. 

771 Requires ``graphviz`` to be installed. 

772 

773 Parameters 

774 ---------- 

775 filename : str or None, optional 

776 The name of the file to write to disk. If the provided `filename` 

777 doesn't include an extension, '.png' will be used by default. 

778 If `filename` is None, no file will be written, and the graph is 

779 rendered in the Jupyter notebook only. 

780 format : {'png', 'pdf', 'dot', 'svg', 'jpeg', 'jpg'}, optional 

781 Format in which to write output file. Default is 'svg'. 

782 **kwargs 

783 Additional keyword arguments to forward to ``to_graphviz``. 

784 """ 

785 from dask.dot import graphviz_to_file 

786 

787 g = self._to_graphviz(**kwargs) 

788 graphviz_to_file(g, filename, format) 

789 return g 

790 

791 def walk(self) -> Generator[Expr]: 

792 """Iterate through all expressions in the tree 

793 

794 Returns 

795 ------- 

796 nodes 

797 Generator of Expr instances in the graph. 

798 Ordering is a depth-first search of the expression tree 

799 """ 

800 stack = [self] 

801 seen = set() 

802 while stack: 

803 node = stack.pop() 

804 if node._name in seen: 

805 continue 

806 seen.add(node._name) 

807 

808 for dep in node.dependencies(): 

809 stack.append(dep) 

810 

811 yield node 

812 

813 def find_operations(self, operation: type | tuple[type]) -> Generator[Expr]: 

814 """Search the expression graph for a specific operation type 

815 

816 Parameters 

817 ---------- 

818 operation 

819 The operation type to search for. 

820 

821 Returns 

822 ------- 

823 nodes 

824 Generator of `operation` instances. Ordering corresponds 

825 to a depth-first search of the expression graph. 

826 """ 

827 assert ( 

828 isinstance(operation, tuple) 

829 and all(issubclass(e, Expr) for e in operation) 

830 or issubclass(operation, Expr) # type: ignore[arg-type] 

831 ), "`operation` must be an `Expr` subclass)" 

832 return (expr for expr in self.walk() if isinstance(expr, operation)) 

833 

834 def __getattr__(self, key): 

835 try: 

836 return object.__getattribute__(self, key) 

837 except AttributeError as err: 

838 if key.startswith("_meta"): 

839 # Avoid a recursive loop if/when `self._meta*` 

840 # produces an `AttributeError` 

841 raise RuntimeError( 

842 f"Failed to generate metadata for {self}. " 

843 "This operation may not be supported by the current backend." 

844 ) 

845 

846 # Allow operands to be accessed as attributes 

847 # as long as the keys are not already reserved 

848 # by existing methods/properties 

849 _parameters = type(self)._parameters 

850 if key in _parameters: 

851 idx = _parameters.index(key) 

852 return self.operands[idx] 

853 

854 raise AttributeError( 

855 f"{err}\n\n" 

856 "This often means that you are attempting to use an unsupported " 

857 f"API function.." 

858 ) 

859 

860 

861class SingletonExpr(Expr): 

862 """A singleton Expr class 

863 

864 This is used to treat the subclassed expression as a singleton. Singletons 

865 are deduplicated by expr._name which is typically based on the dask.tokenize 

866 output. 

867 

868 This is a crucial performance optimization for expressions that walk through 

869 an optimizer and are recreated repeatedly but isn't safe for objects that 

870 cannot be reliably or quickly tokenized. 

871 """ 

872 

873 _instances: weakref.WeakValueDictionary[str, SingletonExpr] 

874 

875 def __new__(cls, *args, _determ_token=None, **kwargs): 

876 if not hasattr(cls, "_instances"): 

877 cls._instances = weakref.WeakValueDictionary() 

878 inst = super().__new__(cls, *args, _determ_token=_determ_token, **kwargs) 

879 _name = inst._name 

880 if _name in cls._instances and cls.__init__ == object.__init__: 

881 return cls._instances[_name] 

882 

883 cls._instances[_name] = inst 

884 return inst 

885 

886 

887def collect_dependents(expr) -> defaultdict: 

888 dependents = defaultdict(list) 

889 stack = [expr] 

890 seen = set() 

891 while stack: 

892 node = stack.pop() 

893 if node._name in seen: 

894 continue 

895 seen.add(node._name) 

896 

897 for dep in node.dependencies(): 

898 stack.append(dep) 

899 dependents[dep._name].append(weakref.ref(node)) 

900 return dependents 

901 

902 

903def optimize(expr: Expr, fuse: bool = True) -> Expr: 

904 """High level query optimization 

905 

906 This leverages three optimization passes: 

907 

908 1. Class based simplification using the ``_simplify`` function and methods 

909 2. Blockwise fusion 

910 

911 Parameters 

912 ---------- 

913 expr: 

914 Input expression to optimize 

915 fuse: 

916 whether or not to turn on blockwise fusion 

917 

918 See Also 

919 -------- 

920 simplify 

921 optimize_blockwise_fusion 

922 """ 

923 stage: OptimizerStage = "fused" if fuse else "simplified-physical" 

924 

925 return optimize_until(expr, stage) 

926 

927 

928def optimize_until(expr: Expr, stage: OptimizerStage) -> Expr: 

929 result = expr 

930 if stage == "logical": 

931 return result 

932 

933 # Simplify 

934 expr = result.simplify() 

935 if stage == "simplified-logical": 

936 return expr 

937 

938 # Manipulate Expression to make it more efficient 

939 if dask.config.get("optimization.tune.active", True): 

940 expr = expr.rewrite(kind="tune", rewritten={}) 

941 if stage == "tuned-logical": 

942 return expr 

943 

944 # Lower 

945 expr = expr.lower_completely() 

946 if stage == "physical": 

947 return expr 

948 

949 # Simplify again 

950 expr = expr.simplify() 

951 if stage == "simplified-physical": 

952 return expr 

953 

954 # Final graph-specific optimizations 

955 expr = expr.fuse() 

956 if stage == "fused": 

957 return expr 

958 

959 raise ValueError(f"Stage {stage!r} not supported.") 

960 

961 

962class LLGExpr(Expr): 

963 """Low Level Graph Expression""" 

964 

965 _parameters = ["dsk"] 

966 

967 def __dask_keys__(self): 

968 return list(self.operand("dsk")) 

969 

970 def _layer(self) -> dict: 

971 return ensure_dict(self.operand("dsk")) 

972 

973 

974class HLGExpr(Expr): 

975 _parameters = [ 

976 "dsk", 

977 "low_level_optimizer", 

978 "output_keys", 

979 "postcompute", 

980 "_cached_optimized", 

981 ] 

982 _defaults = { 

983 "low_level_optimizer": None, 

984 "output_keys": None, 

985 "postcompute": None, 

986 "_cached_optimized": None, 

987 } 

988 

989 @property 

990 def hlg(self): 

991 return self.operand("dsk") 

992 

993 @staticmethod 

994 def from_collection(collection, optimize_graph=True): 

995 from dask.highlevelgraph import HighLevelGraph 

996 

997 if hasattr(collection, "dask"): 

998 dsk = collection.dask.copy() 

999 else: 

1000 dsk = collection.__dask_graph__() 

1001 

1002 # Delayed objects still ship with low level graphs as `dask` when going 

1003 # through optimize / persist 

1004 if not isinstance(dsk, HighLevelGraph): 

1005 

1006 dsk = HighLevelGraph.from_collections( 

1007 str(id(collection)), dsk, dependencies=() 

1008 ) 

1009 if optimize_graph and not hasattr(collection, "__dask_optimize__"): 

1010 warnings.warn( 

1011 f"Collection {type(collection)} does not define a " 

1012 "`__dask_optimize__` method. In the future this will raise. " 

1013 "If no optimization is desired, please set this to `None`.", 

1014 PendingDeprecationWarning, 

1015 ) 

1016 low_level_optimizer = None 

1017 else: 

1018 low_level_optimizer = ( 

1019 collection.__dask_optimize__ if optimize_graph else None 

1020 ) 

1021 return HLGExpr( 

1022 dsk=dsk, 

1023 low_level_optimizer=low_level_optimizer, 

1024 output_keys=collection.__dask_keys__(), 

1025 postcompute=collection.__dask_postcompute__(), 

1026 ) 

1027 

1028 def finalize_compute(self): 

1029 return HLGFinalizeCompute( 

1030 self, 

1031 low_level_optimizer=self.low_level_optimizer, 

1032 output_keys=self.output_keys, 

1033 postcompute=self.postcompute, 

1034 ) 

1035 

1036 def __dask_annotations__(self) -> dict[str, dict[Key, object]]: 

1037 # optimization has to be called (and cached) since blockwise fusion can 

1038 # alter annotations 

1039 # see `dask.blockwise.(_fuse_annotations|_can_fuse_annotations)` 

1040 dsk = self._optimized_dsk 

1041 annotations_by_type: defaultdict[str, dict[Key, object]] = defaultdict(dict) 

1042 for layer in dsk.layers.values(): 

1043 if layer.annotations: 

1044 annot = layer.annotations 

1045 for annot_type, value in annot.items(): 

1046 annotations_by_type[annot_type].update( 

1047 {k: (value(k) if callable(value) else value) for k in layer} 

1048 ) 

1049 return dict(annotations_by_type) 

1050 

1051 def __dask_keys__(self): 

1052 if (keys := self.operand("output_keys")) is not None: 

1053 return keys 

1054 dsk = self.hlg 

1055 # Note: This will materialize 

1056 dependencies = dsk.get_all_dependencies() 

1057 leafs = set(dependencies) 

1058 for val in dependencies.values(): 

1059 leafs -= val 

1060 self.output_keys = list(leafs) 

1061 return self.output_keys 

1062 

1063 @functools.cached_property 

1064 def _optimized_dsk(self) -> HighLevelGraph: 

1065 from dask.highlevelgraph import HighLevelGraph 

1066 

1067 optimizer = self.low_level_optimizer 

1068 keys = self.__dask_keys__() 

1069 dsk = self.hlg 

1070 if (optimizer := self.low_level_optimizer) is not None: 

1071 dsk = optimizer(dsk, keys) 

1072 return HighLevelGraph.merge(dsk) 

1073 

1074 @property 

1075 def deterministic_token(self): 

1076 if not self._determ_token: 

1077 self._determ_token = uuid.uuid4().hex 

1078 return self._determ_token 

1079 

1080 def _layer(self) -> dict: 

1081 dsk = self._optimized_dsk 

1082 return ensure_dict(dsk) 

1083 

1084 

1085class _HLGExprGroup(HLGExpr): 

1086 # Identical to HLGExpr 

1087 # Used internally to determine how output keys are supposed to be returned 

1088 pass 

1089 

1090 

1091class _HLGExprSequence(Expr): 

1092 

1093 def __getitem__(self, other): 

1094 return self.operands[other] 

1095 

1096 def _operands_for_repr(self): 

1097 return [ 

1098 f"name={self.operand('name')!r}", 

1099 f"dsk={self.operand('dsk')!r}", 

1100 ] 

1101 

1102 def _tree_repr_lines(self, indent=0, recursive=True): 

1103 return self._operands_for_repr() 

1104 

1105 def finalize_compute(self): 

1106 return _HLGExprSequence(*[op.finalize_compute() for op in self.operands]) 

1107 

1108 def _tune_down(self): 

1109 if len(self.operands) == 1: 

1110 return None 

1111 from dask.highlevelgraph import HighLevelGraph 

1112 

1113 groups = toolz.groupby( 

1114 lambda x: x.low_level_optimizer if isinstance(x, HLGExpr) else None, 

1115 self.operands, 

1116 ) 

1117 exprs = [] 

1118 changed = False 

1119 for optimizer, group in groups.items(): 

1120 if len(group) > 1: 

1121 graphs = [expr.hlg for expr in group] 

1122 

1123 changed = True 

1124 dsk = HighLevelGraph.merge(*graphs) 

1125 hlg_group = _HLGExprGroup( 

1126 dsk=dsk, 

1127 low_level_optimizer=optimizer, 

1128 output_keys=[v.__dask_keys__() for v in group], 

1129 postcompute=[g.postcompute for g in group], 

1130 ) 

1131 exprs.append(hlg_group) 

1132 else: 

1133 exprs.append(group[0]) 

1134 if not changed: 

1135 return None 

1136 return _HLGExprSequence(*exprs) 

1137 

1138 @functools.cached_property 

1139 def _optimized_dsk(self) -> HighLevelGraph: 

1140 from dask.highlevelgraph import HighLevelGraph 

1141 

1142 hlgexpr: HLGExpr 

1143 graphs = [] 

1144 # simplify_down ensure there are only one HLGExpr per optimizer/finalizer 

1145 for hlgexpr in self.operands: 

1146 keys = hlgexpr.__dask_keys__() 

1147 dsk = hlgexpr.hlg 

1148 if (optimizer := hlgexpr.low_level_optimizer) is not None: 

1149 dsk = optimizer(dsk, keys) 

1150 graphs.append(dsk) 

1151 

1152 return HighLevelGraph.merge(*graphs) 

1153 

1154 def __dask_graph__(self): 

1155 # This class has to override this and not just _layer to ensure the HLGs 

1156 # are not optimized individually 

1157 return ensure_dict(self._optimized_dsk) 

1158 

1159 _layer = __dask_graph__ 

1160 

1161 def __dask_annotations__(self) -> dict[str, dict[Key, object]]: 

1162 # optimization has to be called (and cached) since blockwise fusion can 

1163 # alter annotations 

1164 # see `dask.blockwise.(_fuse_annotations|_can_fuse_annotations)` 

1165 dsk = self._optimized_dsk 

1166 annotations_by_type: defaultdict[str, dict[Key, object]] = defaultdict(dict) 

1167 for layer in dsk.layers.values(): 

1168 if layer.annotations: 

1169 annot = layer.annotations 

1170 for annot_type, value in annot.items(): 

1171 annots = list( 

1172 (k, (value(k) if callable(value) else value)) for k in layer 

1173 ) 

1174 annotations_by_type[annot_type].update( 

1175 { 

1176 k: v 

1177 for k, v in annots 

1178 if not isinstance(v, _AnnotationsTombstone) 

1179 } 

1180 ) 

1181 if not annotations_by_type[annot_type]: 

1182 del annotations_by_type[annot_type] 

1183 return dict(annotations_by_type) 

1184 

1185 def __dask_keys__(self) -> list: 

1186 all_keys = [] 

1187 for op in self.operands: 

1188 if isinstance(op, _HLGExprGroup): 

1189 all_keys.extend(op.__dask_keys__()) 

1190 else: 

1191 all_keys.append(op.__dask_keys__()) 

1192 return all_keys 

1193 

1194 

1195class _ExprSequence(Expr): 

1196 """A sequence of expressions 

1197 

1198 This is used to be able to optimize multiple collections combined, e.g. when 

1199 being computed simultaneously with ``dask.compute((Expr1, Expr2))``. 

1200 """ 

1201 

1202 def __getitem__(self, other): 

1203 return self.operands[other] 

1204 

1205 def _layer(self) -> dict: 

1206 return toolz.merge(op._layer() for op in self.operands) 

1207 

1208 def __dask_keys__(self) -> list: 

1209 all_keys = [] 

1210 for op in self.operands: 

1211 all_keys.append(list(op.__dask_keys__())) 

1212 return all_keys 

1213 

1214 def __repr__(self): 

1215 return "ExprSequence(" + ", ".join(map(repr, self.operands)) + ")" 

1216 

1217 __str__ = __repr__ 

1218 

1219 def finalize_compute(self): 

1220 return _ExprSequence( 

1221 *(op.finalize_compute() for op in self.operands), 

1222 ) 

1223 

1224 def __dask_annotations__(self): 

1225 annotations_by_type = {} 

1226 for op in self.operands: 

1227 for k, v in op.__dask_annotations__().items(): 

1228 annotations_by_type.setdefault(k, {}).update(v) 

1229 return annotations_by_type 

1230 

1231 def __len__(self): 

1232 return len(self.operands) 

1233 

1234 def __iter__(self): 

1235 return iter(self.operands) 

1236 

1237 def _simplify_down(self): 

1238 from dask.highlevelgraph import HighLevelGraph 

1239 

1240 issue_warning = False 

1241 hlgs = [] 

1242 if any( 

1243 isinstance(op, (HLGExpr, HLGFinalizeCompute, dict)) for op in self.operands 

1244 ): 

1245 for op in self.operands: 

1246 if isinstance(op, (HLGExpr, HLGFinalizeCompute)): 

1247 hlgs.append(op) 

1248 elif isinstance(op, dict): 

1249 hlgs.append( 

1250 HLGExpr( 

1251 dsk=HighLevelGraph.from_collections( 

1252 str(id(op)), op, dependencies=() 

1253 ) 

1254 ) 

1255 ) 

1256 else: 

1257 issue_warning = True 

1258 opt = op.optimize() 

1259 hlgs.append( 

1260 HLGExpr( 

1261 dsk=HighLevelGraph.from_collections( 

1262 opt._name, opt.__dask_graph__(), dependencies=() 

1263 ) 

1264 ) 

1265 ) 

1266 if issue_warning: 

1267 warnings.warn( 

1268 "Computing mixed collections that are backed by " 

1269 "HighlevelGraphs/dicts and Expressions. " 

1270 "This forces Expressions to be materialized. " 

1271 "It is recommended to use only one type and separate the dask." 

1272 "compute calls if necessary.", 

1273 UserWarning, 

1274 ) 

1275 if not hlgs: 

1276 return None 

1277 return _HLGExprSequence(*hlgs) 

1278 

1279 

1280class CompositeExpr(Expr): 

1281 """Private expression grouping many child expressions into one collection.""" 

1282 

1283 _parameters = ["collection"] 

1284 

1285 @property 

1286 def collection(self): 

1287 return self.operand("collection") 

1288 

1289 @property 

1290 def exprs(self): 

1291 return tuple(self.operands[1:]) 

1292 

1293 def dependencies(self): 

1294 return self.exprs 

1295 

1296 def _layer(self) -> dict: 

1297 return {} 

1298 

1299 def __dask_keys__(self) -> list: 

1300 return [expr.__dask_keys__() for expr in self.exprs] 

1301 

1302 def _operands_for_repr(self): 

1303 return [ 

1304 f"collection={type(self.collection).__name__}", 

1305 f"nexprs={len(self.exprs)}", 

1306 ] 

1307 

1308 def finalize_compute(self): 

1309 return CompositeFinalizeCompute(self.collection, *self.exprs) 

1310 

1311 

1312class _AnnotationsTombstone: ... 

1313 

1314 

1315class FinalizeCompute(Expr): 

1316 _parameters = ["expr"] 

1317 

1318 def _simplify_down(self): 

1319 return self.expr.finalize_compute() 

1320 

1321 

1322def _convert_dask_keys(keys): 

1323 from dask._task_spec import List, TaskRef 

1324 

1325 assert isinstance(keys, list) 

1326 new_keys = [] 

1327 for key in keys: 

1328 if isinstance(key, list): 

1329 new_keys.append(_convert_dask_keys(key)) 

1330 else: 

1331 new_keys.append(TaskRef(key)) 

1332 return List(*new_keys) 

1333 

1334 

1335class HLGFinalizeCompute(HLGExpr): 

1336 

1337 def _simplify_down(self): 

1338 if not self.postcompute: 

1339 return self.dsk 

1340 

1341 from dask.delayed import Delayed 

1342 

1343 # Skip finalization for Delayed 

1344 if self.dsk.postcompute == Delayed.__dask_postcompute__(self.dsk): 

1345 return self.dsk 

1346 return self 

1347 

1348 @property 

1349 def _name(self): 

1350 return f"finalize-{super()._name}" 

1351 

1352 def __dask_graph__(self): 

1353 # The base class __dask_graph__ will not just materialize this layer but 

1354 # also that of its dependencies, i.e. it will render the finalized and 

1355 # the non-finalized graph and combine them. We only want the finalized 

1356 # so we're overriding this. 

1357 # This is an artifact generated because the wrapped expression is 

1358 # identified automatically as a dependency but HLG expressions are not 

1359 # working in this layered way. 

1360 return self._layer() 

1361 

1362 @property 

1363 def hlg(self): 

1364 expr = self.operand("dsk") 

1365 layers = expr.dsk.layers.copy() 

1366 deps = expr.dsk.dependencies.copy() 

1367 keys = expr.__dask_keys__() 

1368 if isinstance(expr.postcompute, list): 

1369 postcomputes = expr.postcompute 

1370 else: 

1371 postcomputes = [expr.postcompute] 

1372 tasks = [ 

1373 Task(self._name, func, _convert_dask_keys(keys), *extra_args) 

1374 for func, extra_args in postcomputes 

1375 ] 

1376 from dask.highlevelgraph import HighLevelGraph, MaterializedLayer 

1377 

1378 leafs = set(deps) 

1379 for val in deps.values(): 

1380 leafs -= val 

1381 for t in tasks: 

1382 layers[t.key] = MaterializedLayer({t.key: t}) 

1383 deps[t.key] = leafs 

1384 return HighLevelGraph(layers, dependencies=deps) 

1385 

1386 def __dask_keys__(self): 

1387 return [self._name] 

1388 

1389 

1390class CompositeFinalizeCompute(CompositeExpr): 

1391 def __dask_keys__(self): 

1392 return [self._name] 

1393 

1394 def _layer(self) -> dict: 

1395 func, extra_args = self.collection.__dask_postcompute__() 

1396 keys = [expr.__dask_keys__() for expr in self.exprs] 

1397 return { 

1398 self._name: Task( 

1399 self._name, 

1400 func, 

1401 _convert_dask_keys(keys), 

1402 *extra_args, 

1403 ) 

1404 } 

1405 

1406 

1407class ProhibitReuse(Expr): 

1408 """ 

1409 An expression that guarantees that all keys are suffixes with a unique id. 

1410 This can be used to break a common subexpression apart. 

1411 """ 

1412 

1413 _parameters = ["expr"] 

1414 _ALLOWED_TYPES = [HLGExpr, LLGExpr, HLGFinalizeCompute, _HLGExprSequence] 

1415 

1416 def __dask_keys__(self): 

1417 return self._modify_keys(self.expr.__dask_keys__()) 

1418 

1419 @staticmethod 

1420 def _identity(obj): 

1421 return obj 

1422 

1423 @functools.cached_property 

1424 def _suffix(self): 

1425 return uuid.uuid4().hex 

1426 

1427 def _modify_keys(self, k): 

1428 if isinstance(k, list): 

1429 return [self._modify_keys(kk) for kk in k] 

1430 elif isinstance(k, tuple): 

1431 return (self._modify_keys(k[0]),) + k[1:] 

1432 elif isinstance(k, (int, float)): 

1433 k = str(k) 

1434 return f"{k}-{self._suffix}" 

1435 

1436 def _simplify_down(self): 

1437 # FIXME: Shuffling cannot be rewritten since the barrier key is 

1438 # hardcoded. Skipping this here should do the trick most of the time 

1439 if not isinstance( 

1440 self.expr, 

1441 tuple(self._ALLOWED_TYPES), 

1442 ): 

1443 return self.expr 

1444 

1445 def __dask_graph__(self): 

1446 try: 

1447 from distributed.shuffle._core import P2PBarrierTask 

1448 except ModuleNotFoundError: 

1449 P2PBarrierTask = type(None) 

1450 dsk = convert_legacy_graph(self.expr.__dask_graph__()) 

1451 

1452 subs = {old_key: self._modify_keys(old_key) for old_key in dsk} 

1453 dsk2 = {} 

1454 for old_key, new_key in subs.items(): 

1455 t = dsk[old_key] 

1456 if isinstance(t, P2PBarrierTask): 

1457 warnings.warn( 

1458 "Cannot block reusing for graphs including a " 

1459 "P2PBarrierTask. This may cause unexpected results. " 

1460 "This typically happens when converting a dask " 

1461 "DataFrame to delayed objects.", 

1462 UserWarning, 

1463 ) 

1464 return dsk 

1465 dsk2[new_key] = Task( 

1466 new_key, 

1467 ProhibitReuse._identity, 

1468 t.substitute(subs), 

1469 ) 

1470 

1471 dsk2.update(dsk) 

1472 return dsk2 

1473 

1474 _layer = __dask_graph__