diff --git a/integration_tests/old_benchmarks/anemoi.py b/integration_tests/old_benchmarks/anemoi.py index c123acb5..a56f7381 100644 --- a/integration_tests/old_benchmarks/anemoi.py +++ b/integration_tests/old_benchmarks/anemoi.py @@ -6,7 +6,7 @@ # granted to it by virtue of its status as an intergovernmental organisation # nor does it submit to any jurisdiction. -from earthkit.workflows import Cascade +from earthkit.workflows.visualise import visualise def get_graph(lead_time, ensemble_members, CKPT=None, date="2024-12-02T00:00"): @@ -18,7 +18,7 @@ def get_graph(lead_time, ensemble_members, CKPT=None, date="2024-12-02T00:00"): result = model_action.mean(dim="ensemble_member") result = result.map(print) - cascade_model = Cascade.from_actions([result.sel(param="2t")]) + cascade_model = result.sel(param="2t").graph() - cascade_model.visualise("model_running.html", preset="blob", cdn_resources="in_line") - return cascade_model._graph + visualise(cascade_model, "model_running.html", preset="blob", cdn_resources="in_line") + return cascade_model diff --git a/src/earthkit/workflows/__init__.py b/src/earthkit/workflows/__init__.py index 7f211331..75608eff 100644 --- a/src/earthkit/workflows/__init__.py +++ b/src/earthkit/workflows/__init__.py @@ -16,41 +16,8 @@ # assuming editable install etc pass from . import fluent, mark -from .graph import Graph, deduplicate_nodes -from .graph.export import deserialise, serialise - - -class Cascade: - def __init__(self, graph: Graph = Graph([])): - self._graph = graph - - @classmethod - def from_actions(cls, actions): - graph = Graph([]) - for action in actions: - graph += action.graph() - return cls(deduplicate_nodes(graph)) - - def visualise(self, *args, **kwargs): - from .visualise import visualise as _visualise_fn - - return _visualise_fn(self._graph, *args, **kwargs) - - def __add__(self, other: "Cascade") -> "Cascade": - if not isinstance(other, Cascade): - return NotImplemented - return Cascade(deduplicate_nodes(self._graph + other._graph)) - - def __iadd__(self, other: "Cascade") -> "Cascade": - if not isinstance(other, Cascade): - return NotImplemented - self._graph += other._graph - self._graph = deduplicate_nodes(self._graph) - return self - __all__ = [ "mark", "fluent", - "Cascade", ]