1- """DAGAnnotator and Composer implementation."""
1+ """Composer implementation."""
22
33import collections
4- from dataclasses import dataclass
5- from pathlib import Path
64import functools
75import json
86import os
97from concurrent .futures import ThreadPoolExecutor
8+ from dataclasses import dataclass
109from itertools import product
10+ from pathlib import Path
1111from typing import Any , Optional
1212
1313import pandas as pd
14+ from modelbench .cache import DiskCache , NullCache
1415from tqdm import tqdm
1516
16- from modelbench .cache import DiskCache , NullCache
1717from modelplane .evaluator .context import EvalContext , NodeOutput
1818from modelplane .evaluator .cost import CostInfo , RealizedCost
1919from modelplane .evaluator .nodes import Arbiter , CacheableNodeMixin , ComposerNode , Gate
@@ -30,12 +30,21 @@ def wrapper(self, *args, **kwargs):
3030
3131
3232@dataclass
33- class DAGOutput :
34- verdict : Verdict
33+ class _DAGOutput :
3534 node_outputs : dict [str , NodeOutput ]
3635 total_cost : RealizedCost
3736
3837
38+ @dataclass
39+ class SuccessfulDAGOutput (_DAGOutput ):
40+ verdict : Verdict
41+
42+
43+ @dataclass
44+ class FailedDAGOutput (_DAGOutput ):
45+ error : Exception
46+
47+
3948class Composer :
4049 """DAG of ComposerNodes.
4150
@@ -58,7 +67,9 @@ class Composer:
5867 results_df = dag.run_dataframe(df)
5968 """
6069
61- def __init__ (self , name : str , verdict_type : type , cache_path : Optional [Path ] = None ) -> None :
70+ def __init__ (
71+ self , name : str , verdict_type : type , cache_path : Optional [Path ] = None
72+ ) -> None :
6273 self .name = name
6374 self ._nodes : dict [str , ComposerNode ] = {}
6475 self ._root_nodes : list [str ] = []
@@ -79,6 +90,10 @@ def verdict_type(self) -> type:
7990 def df_output_col (self ) -> str :
8091 return f"{ self .name } _output"
8192
93+ @property
94+ def df_error_col (self ) -> str :
95+ return f"{ self .name } _error"
96+
8297 @property
8398 def df_dag_run_col (self ) -> str :
8499 return f"{ self .name } _dag_run"
@@ -100,7 +115,11 @@ def add_node(
100115 self ._nodes [node .name ] = node
101116 self ._validated = False
102117 if isinstance (node , CacheableNodeMixin ):
103- self ._node_caches [node .name ] = DiskCache (self ._cache_path / node .name ) if self ._cache_path else NullCache ()
118+ self ._node_caches [node .name ] = (
119+ DiskCache (self ._cache_path / node .name )
120+ if self ._cache_path
121+ else NullCache ()
122+ )
104123 return self
105124
106125 def _validate_and_build (self ) -> None :
@@ -172,7 +191,21 @@ def _validate_and_build(self) -> None:
172191 self ._root_nodes = root_nodes
173192 self ._ordered = ordered
174193
175- def _run_traced (self , ctx : EvalContext ) -> tuple [DAGOutput , set [tuple [str , str ]]]:
194+ def _run_node (self , node : ComposerNode , ctx : EvalContext ) -> NodeOutput :
195+ if isinstance (node , CacheableNodeMixin ):
196+ key = node .cache_key (ctx )
197+ if key in self ._node_caches [node .name ]:
198+ return self ._node_caches [node .name ][key ]
199+ else :
200+ output = node .run (ctx )
201+ self ._node_caches [node .name ][key ] = output
202+ return output
203+ else :
204+ return node .run (ctx )
205+
206+ def _run_traced (
207+ self , ctx : EvalContext
208+ ) -> tuple [SuccessfulDAGOutput | FailedDAGOutput , set [tuple [str , str ]]]:
176209 """Execute the DAG and return (final verdict, node outputs, realized costs, traversed edges)."""
177210 node_outputs : dict [str , NodeOutput ] = {}
178211 traversed_edges : set [tuple [str , str ]] = set ()
@@ -189,20 +222,20 @@ def _run_traced(self, ctx: EvalContext) -> tuple[DAGOutput, set[tuple[str, str]]
189222 }
190223 )
191224 node = self ._nodes [node_name ]
192- if isinstance ( node , CacheableNodeMixin ) :
193- key = node . cache_key ( ctx )
194- if key in self . _node_caches [ node . name ] :
195- output = self . _node_caches [ node . name ][ key ]
196- else :
197- output = node . run ( ctx )
198- self . _node_caches [ node . name ][ key ] = output
199- else :
200- output = node . run ( ctx )
225+ try :
226+ output = self . _run_node ( node , ctx )
227+ except Exception as e :
228+ return (
229+ FailedDAGOutput (
230+ node_outputs = node_outputs , total_cost = total_cost , error = e
231+ ),
232+ traversed_edges ,
233+ )
201234 node_outputs [node_name ] = output
202235 total_cost += output .realized_cost
203236 if isinstance (output .value , Verdict ):
204237 traversed_edges .add ((node_name , output .value .name ))
205- dag_output = DAGOutput (
238+ dag_output = SuccessfulDAGOutput (
206239 verdict = output .value ,
207240 node_outputs = node_outputs ,
208241 total_cost = total_cost ,
@@ -213,7 +246,7 @@ def _run_traced(self, ctx: EvalContext) -> tuple[DAGOutput, set[tuple[str, str]]
213246 traversed_edges .add ((node_name , t ))
214247 if isinstance (target , Verdict ):
215248 return (
216- DAGOutput (
249+ SuccessfulDAGOutput (
217250 verdict = target ,
218251 node_outputs = node_outputs ,
219252 total_cost = total_cost ,
@@ -224,7 +257,7 @@ def _run_traced(self, ctx: EvalContext) -> tuple[DAGOutput, set[tuple[str, str]]
224257 raise ValueError ("DAG execution completed without reaching a Verdict node." )
225258
226259 @requires_validate_and_build
227- def run (self , ctx : EvalContext ) -> DAGOutput :
260+ def run (self , ctx : EvalContext ) -> SuccessfulDAGOutput | FailedDAGOutput :
228261 """Execute the DAG on a single prompt/response and get the output,
229262 node outputs, and overall realized cost."""
230263 dag_output , _ = self ._run_traced (ctx )
@@ -241,7 +274,7 @@ def run_dataframe(
241274 ) -> pd .DataFrame :
242275 """Run the DAG over every row of a DataFrame."""
243276
244- def _run_row (row : Any ) -> DAGOutput :
277+ def _run_row (row : Any ) -> SuccessfulDAGOutput | FailedDAGOutput :
245278 ctx = EvalContext (
246279 prompt = str (row [prompt_col ]),
247280 response = str (row [response_col ]),
@@ -262,7 +295,14 @@ def _run_row(row: Any) -> DAGOutput:
262295
263296 result_df = pd .DataFrame (
264297 {
265- self .df_output_col : [r .verdict .name for r in records ],
298+ self .df_output_col : [
299+ r .verdict .name if isinstance (r , SuccessfulDAGOutput ) else None
300+ for r in records
301+ ],
302+ self .df_error_col : [
303+ str (r .error ) if isinstance (r , FailedDAGOutput ) else None
304+ for r in records
305+ ],
266306 self .df_dag_run_col : [
267307 json .dumps ({k : v .to_dict () for k , v in r .node_outputs .items ()})
268308 for r in records
@@ -587,6 +627,10 @@ def visualize_run(self, ctx: EvalContext):
587627 return self ._visualize (
588628 node_outputs = dag_output .node_outputs ,
589629 traversed_edges = traversed_edges ,
590- final_output = dag_output .verdict ,
630+ final_output = (
631+ dag_output .verdict
632+ if isinstance (dag_output , SuccessfulDAGOutput )
633+ else None
634+ ),
591635 ctx = ctx ,
592636 )
0 commit comments