Skip to content

Commit 23e7ea7

Browse files
committed
Improved DAG exception handling.
1 parent fdb0224 commit 23e7ea7

5 files changed

Lines changed: 117 additions & 31 deletions

File tree

‎src/modelplane/evaluator/annotator.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
from modelgauge.sut import SUTResponse
55

66
from modelplane.evaluator.context import EvalContext
7-
from modelplane.evaluator.dag import Composer
7+
from modelplane.evaluator.dag import Composer, SuccessfulDAGOutput
88
from modelplane.evaluator.verdict import Verdict
99

1010

@@ -29,4 +29,8 @@ def translate_prompt(
2929
)
3030

3131
def annotate(self, annotation_request: EvalContext) -> Verdict:
32-
return self.dag.run(annotation_request).verdict
32+
dag_output = self.dag.run(annotation_request)
33+
if isinstance(dag_output, SuccessfulDAGOutput):
34+
return dag_output.verdict
35+
else:
36+
raise dag_output.error

‎src/modelplane/evaluator/dag.py‎

Lines changed: 68 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -1,19 +1,19 @@
1-
"""DAGAnnotator and Composer implementation."""
1+
"""Composer implementation."""
22

33
import collections
4-
from dataclasses import dataclass
5-
from pathlib import Path
64
import functools
75
import json
86
import os
97
from concurrent.futures import ThreadPoolExecutor
8+
from dataclasses import dataclass
109
from itertools import product
10+
from pathlib import Path
1111
from typing import Any, Optional
1212

1313
import pandas as pd
14+
from modelbench.cache import DiskCache, NullCache
1415
from tqdm import tqdm
1516

16-
from modelbench.cache import DiskCache, NullCache
1717
from modelplane.evaluator.context import EvalContext, NodeOutput
1818
from modelplane.evaluator.cost import CostInfo, RealizedCost
1919
from 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+
3948
class 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
)

‎src/modelplane/evaluator/safety.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,7 @@ def __init__(self, name: str, annotator: Annotator) -> None:
4848
super().__init__(name=name)
4949
self.annotator = annotator
5050

51-
def run(self, ctx: EvalContext) -> Safety:
51+
def run(self, ctx: EvalContext) -> NodeOutput:
5252
prompt = TextPrompt(text=ctx.prompt)
5353
response = SUTResponse(text=ctx.response)
5454
annotation = self.annotator.process(prompt, response)

‎tests/unit/evaluator/mocks.py‎

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,26 @@
11
from modelplane.evaluator.context import EvalContext, NodeOutput
22
from modelplane.evaluator.cost import CostInfo
3-
from modelplane.evaluator.nodes import Arbiter, CacheableNodeMixin, Enricher, Gate, LLMCostMixin
4-
from modelplane.evaluator.verdict import Verdict
3+
from modelplane.evaluator.nodes import (
4+
Arbiter,
5+
CacheableNodeMixin,
6+
Enricher,
7+
Gate,
8+
LLMCostMixin,
9+
)
510
from modelplane.evaluator.safety import Safety
11+
from modelplane.evaluator.verdict import Verdict
612

713

814
def context_token_count(ctx: EvalContext) -> int:
915
return len(ctx.prompt.split() + ctx.response.split())
1016

1117

18+
class FailingNode(Enricher):
19+
20+
def run(self, ctx: EvalContext) -> NodeOutput:
21+
raise RuntimeError("I'm afraid I can't do that, Dave.")
22+
23+
1224
class PassthroughGate(Gate, LLMCostMixin):
1325
ROUTE_TO_TAKE: bool
1426

‎tests/unit/evaluator/test_safety.py‎

Lines changed: 28 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,15 @@
1+
import pytest
12
from modelgauge.annotation import SafetyAnnotation
23
from modelgauge.annotators.demo_annotator import DemoYBadAnnotator
34
from modelgauge.prompt import TextPrompt
45
from modelgauge.sut import SUTResponse
5-
import pytest
66

7-
from modelplane.evaluator.dag import Composer
7+
from modelplane.evaluator.dag import Composer, FailedDAGOutput
88
from modelplane.evaluator.safety import AnnotatorArbiter, Safety, SafetyDAGAnnotator
99
from modelplane.evaluator.verdict import Verdict
1010

11+
from .mocks import FailingNode
12+
1113

1214
def test_safety_annotator_arbiter(sample_ctx):
1315
annotator = DemoYBadAnnotator("demo_annotator")
@@ -34,3 +36,27 @@ def test_safety_dag_with_bad_verdict_type():
3436
match="All outputs of the DAG must be of type Safety.",
3537
):
3638
SafetyDAGAnnotator("bad_dag", Composer("bad_dag", verdict_type=Verdict))
39+
40+
41+
def test_safety_dag_with_bad_node(sample_ctx, threshold_arbiter):
42+
failing_node = FailingNode(name="failing_node", routes=["threshold_arbiter"])
43+
dag = (
44+
Composer(
45+
"bad_node_dag",
46+
verdict_type=Safety,
47+
)
48+
.add_node(failing_node)
49+
.add_node(threshold_arbiter)
50+
)
51+
dag_output = dag.run(sample_ctx)
52+
assert isinstance(dag_output, FailedDAGOutput)
53+
assert str(dag_output.error) == "I'm afraid I can't do that, Dave."
54+
55+
dag_annotator = SafetyDAGAnnotator("safety_annotator", dag)
56+
with pytest.raises(
57+
type(dag_output.error), match="I'm afraid I can't do that, Dave."
58+
):
59+
dag_annotator.process(
60+
prompt=TextPrompt(text=sample_ctx.prompt),
61+
response=SUTResponse(text=sample_ctx.response),
62+
)

0 commit comments

Comments
 (0)