Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 107 additions & 0 deletions src/furu/_pytree.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
from __future__ import annotations

import copy
import dataclasses
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from typing import Literal, cast

from pydantic import BaseModel

type IsLeaf = Callable[[object], bool]
type NodeKind = Literal[
"list", "tuple", "set", "frozenset", "dict", "dataclass", "pydantic"
]


@dataclass(frozen=True)
class PyTreeNode:
kind: NodeKind
value: object
entries: tuple[object, ...]
children: tuple[object, ...]

def unflatten(self, children: Iterable[object]) -> object:
values = tuple(children)
match self.kind:
case "list" | "set" | "frozenset":
return type(self.value)(values)
case "tuple":
return (
type(self.value)(*values)
if hasattr(type(self.value), "_fields")
else type(self.value)(values)
)
case "dict":
return type(self.value)(zip(self.entries, values, strict=True))
case "dataclass":
result = copy.copy(self.value)
for name, child in zip(self.entries, values, strict=True):
object.__setattr__(result, cast(str, name), child)
return result
case "pydantic":
assert isinstance(self.value, BaseModel)
return self.value.model_copy(
update=dict(
zip(cast(tuple[str, ...], self.entries), values, strict=True)
)
)


def tree_node(value: object) -> PyTreeNode | None:
if isinstance(value, BaseModel):
entries = tuple(type(value).model_fields)
kind: NodeKind = "pydantic"
children = tuple(getattr(value, name) for name in entries)
elif dataclasses.is_dataclass(value) and not isinstance(value, type):
entries = tuple(field.name for field in dataclasses.fields(value))
kind = "dataclass"
children = tuple(getattr(value, name) for name in entries)
elif isinstance(value, (list, tuple)):
children = tuple(value)
entries = tuple(range(len(children)))
kind = "list" if isinstance(value, list) else "tuple"
elif isinstance(value, (set, frozenset)):
children = tuple(
sorted(
value,
key=lambda item: (
type(item).__module__, type(item).__qualname__, repr(item)
),
)
)
entries = tuple(range(len(children)))
kind = "frozenset" if isinstance(value, frozenset) else "set"
elif isinstance(value, dict):
mapping = cast(dict[object, object], value)
entries = tuple(mapping)
children = tuple(mapping[key] for key in entries)
kind = "dict"
else:
return None
return PyTreeNode(kind, value, entries, children)


def tree_map(
function: Callable[[object], object],
value: object,
*,
is_leaf: IsLeaf | None = None,
) -> object:
node = None if is_leaf is not None and is_leaf(value) else tree_node(value)
if node is None:
return function(value)
return node.unflatten(
tree_map(function, child, is_leaf=is_leaf) for child in node.children
)


def tree_leaves(value: object, *, is_leaf: IsLeaf | None = None) -> list[object]:
leaves: list[object] = []

def collect(leaf: object) -> object:
leaves.append(leaf)
return leaf

tree_map(collect, value, is_leaf=is_leaf)
return leaves
14 changes: 8 additions & 6 deletions src/furu/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,9 @@
from collections.abc import Sequence
from dataclasses import dataclass, field
from functools import cached_property
from typing import TYPE_CHECKING, assert_never
from typing import TYPE_CHECKING, assert_never, cast

from furu._pytree import tree_leaves
from furu.core import Spec
from furu.dependencies import collect_declared_refs
from furu.metadata import ArtifactSpec
Expand All @@ -27,11 +28,11 @@ def batch_group(self) -> tuple[object, int] | None:
return _batch_group(self.obj)


def _add_to_dag(coordinator: ExecutionCoordinator, objs: Sequence[Spec]) -> None:
if any(not isinstance(obj, Spec) for obj in objs):
# TODO: accept pytrees of Spec objects (e.g. nested lists/dicts/dataclasses)
# and flatten them before walking dependencies.
raise TypeError("expected Spec objects")
def _add_to_dag(coordinator: ExecutionCoordinator, obj_tree: object) -> list[Spec]:
leaves = tree_leaves(obj_tree, is_leaf=lambda value: isinstance(value, Spec))
if any(not isinstance(obj, Spec) for obj in leaves):
raise TypeError("expected Spec objects in every PyTree leaf")
objs = cast(list[Spec], leaves)

refs_by_id: dict[str, tuple[Spec, ...]] = {}
newly_added: list[DagNode] = []
Expand Down Expand Up @@ -75,6 +76,7 @@ def _add_to_dag(coordinator: ExecutionCoordinator, objs: Sequence[Spec]) -> None
coordinator.blocked[node.obj.object_id] = node
else:
coordinator.ready[node.obj.object_id] = node
return objs


def _update_dag_blocking_dependencies(
Expand Down
24 changes: 7 additions & 17 deletions src/furu/dependencies.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,11 @@
from collections.abc import Callable, Iterator
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import fields, is_dataclass
from dataclasses import fields
from functools import cached_property
from typing import TYPE_CHECKING, Any, Self, overload

from pydantic import BaseModel as PydanticBaseModel
from furu._pytree import tree_leaves

if TYPE_CHECKING:
from furu.core import Spec
Expand Down Expand Up @@ -47,21 +47,11 @@ def dependency[TSpec: Spec[Any], T](
def find_nested_furu_objects(value: object) -> Iterator[Spec]:
from furu.core import Spec

match value:
case Spec():
yield value
case _ if is_dataclass(value) and not isinstance(value, type):
for field in fields(value):
yield from find_nested_furu_objects(getattr(value, field.name))
case PydanticBaseModel():
for name in type(value).model_fields:
yield from find_nested_furu_objects(getattr(value, name))
case tuple() | list() | set() | frozenset():
for item in value:
yield from find_nested_furu_objects(item)
case dict():
for item in value.values():
yield from find_nested_furu_objects(item)
yield from (
leaf
for leaf in tree_leaves(value, is_leaf=lambda item: isinstance(item, Spec))
if isinstance(leaf, Spec)
)


def collect_declared_refs(obj: Spec) -> tuple[Spec, ...]:
Expand Down
10 changes: 5 additions & 5 deletions src/furu/execution/execution_coordinator.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
import hashlib
import threading
import time
from collections.abc import Iterator, Sequence
from collections.abc import Iterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass, field
Expand Down Expand Up @@ -89,9 +89,9 @@ def _counts_detail(self) -> dict[str, object]:
}

@classmethod
def run[ObjsT: Sequence[Spec]](
def run[ObjsT](
cls,
objs: ObjsT, # TODO: support pytrees
objs: ObjsT,
*,
max_retries_per_object: int | None = None,
worker_backends: tuple[WorkerBackend, ...],
Expand All @@ -100,9 +100,9 @@ def run[ObjsT: Sequence[Spec]](
if max_retries_per_object is None:
max_retries_per_object = get_config().worker.max_retries_per_object
coordinator = cls(max_retries_per_object=max_retries_per_object)
_add_to_dag(coordinator, objs)
root_objs = _add_to_dag(coordinator, objs)
digest = hashlib.blake2s(digest_size=16)
for obj in objs:
for obj in root_objs:
digest.update(obj.object_id.encode("utf-8"))
digest.update(b"\0")
coordinator.executor_id = digest.hexdigest()
Expand Down
Loading
Loading