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
75 changes: 73 additions & 2 deletions src/jobflow/core/job.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,23 @@
logger = logging.getLogger(__name__)


def _replace_job_or_flow_with_output(value):
"""Replace Jobs and Flows in nested containers with their outputs."""
from jobflow.core.flow import Flow

if isinstance(value, (Job, Flow)):
return value.output
if isinstance(value, list):
return [_replace_job_or_flow_with_output(item) for item in value]
if isinstance(value, tuple):
return tuple(_replace_job_or_flow_with_output(item) for item in value)
if isinstance(value, dict):
return {
key: _replace_job_or_flow_with_output(item) for key, item in value.items()
}
return value


@dataclass
class JobConfig(MSONable):
"""
Expand Down Expand Up @@ -213,6 +230,10 @@ def get_job(*args, **kwargs) -> Job:
f = met
args = args[1:]

if _current_flow_context.get() is not None:
args = _replace_job_or_flow_with_output(args)
kwargs = _replace_job_or_flow_with_output(kwargs)

return Job(
function=f, function_args=args, function_kwargs=kwargs, **job_kwargs
)
Expand Down Expand Up @@ -604,7 +625,7 @@ def run(self, store: jobflow.JobStore, job_dir: Path = None) -> Response:
from datetime import datetime

from jobflow import CURRENT_JOB
from jobflow.core.flow import get_flow
from jobflow.core.flow import flow_build_context, get_flow
from jobflow.core.schemas import JobStoreDocument

index_str = f", {self.index}" if self.index != 1 else ""
Expand All @@ -626,12 +647,17 @@ def run(self, store: jobflow.JobStore, job_dir: Path = None) -> Response:
if bound is not None and not isinstance(bound, types.ModuleType):
function = types.MethodType(function, bound)

response = function(*self.function_args, **self.function_kwargs)
dynamic_children = []
with flow_build_context(dynamic_children):
response = function(*self.function_args, **self.function_kwargs)
response = Response.from_job_returns(
response, self.output_schema, job_dir=job_dir
)

if response.replace is not None:
response.replace = _expand_dynamic_dependencies(
response.replace, dynamic_children
)
response.replace = prepare_replace(response.replace, self)

if response.addition is not None:
Expand Down Expand Up @@ -1481,6 +1507,51 @@ def prepare_replace(
return replace


def _expand_dynamic_dependencies(replace, dynamic_children):
"""Add dynamically created dependencies required by a replacement."""
from jobflow.core.flow import Flow

candidates = [child for child in dynamic_children if child.host is None]
if not candidates:
return replace

if isinstance(replace, dict):
explicit_jobs = list(replace.values())
output = {key: child.output for key, child in replace.items()}
elif isinstance(replace, (list, tuple)):
explicit_jobs = list(replace)
output = type(replace)(child.output for child in replace)
else:
explicit_jobs = [replace]
output = replace.output

def all_uuids(child):
if isinstance(child, Flow):
return set(child.all_uuids)
return {child.uuid}

explicit_uuids = set().union(*(all_uuids(child) for child in explicit_jobs))
required_uuids = set(explicit_uuids)

while True:
selected = [
child for child in candidates if all_uuids(child) & required_uuids
]
expanded_uuids = set().union(
required_uuids, *(set(child.graph.nodes) for child in selected)
)
if expanded_uuids == required_uuids:
break
required_uuids = expanded_uuids

selected = [child for child in candidates if all_uuids(child) & required_uuids]
selected_uuids = set().union(*(all_uuids(child) for child in selected))
if selected_uuids.issubset(explicit_uuids):
return replace

return Flow(selected, output=output)


def pass_manager_config(
jobs: Job | jobflow.Flow | list[Job | jobflow.Flow],
manager_config: dict[str, Any],
Expand Down
32 changes: 32 additions & 0 deletions tests/core/test_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -230,6 +230,38 @@ def replace_list_flow():
# currently output schema and metadata ignored for all but the last `store_inputs`


def test_replace_response_collects_dynamic_dependencies():
"""Test replacement jobs include dynamically created dependencies."""
from jobflow import Response, job, run_locally

@job
def add_job(a, b):
return a + b

@job
def multiply_job(value, factor):
return value * factor

@job
def make_dynamic_workflow():
add_job(100, 200)
first = add_job(1, 2)
final = multiply_job(first, 3)
return Response(replace=final)

dynamic_job = make_dynamic_workflow()
responses = run_locally(dynamic_job, ensure_success=True)

outputs = [
response.output
for index_to_response in responses.values()
for response in index_to_response.values()
if isinstance(response.output, int)
]
assert sorted(outputs) == [3, 9]
assert len(responses) == 3


def test_job_config(memory_jobstore):
from jobflow import (
CURRENT_JOB,
Expand Down
Loading