This commit is contained in:
Nuno Campos
2023-10-11 18:32:02 +01:00
parent bd8260760d
commit 851e8f88dd
3 changed files with 651 additions and 0 deletions
+112
View File
@@ -0,0 +1,112 @@
from __future__ import annotations
from langchain.chat_models.openai import ChatOpenAI
from langchain.output_parsers.openai_functions import JsonOutputFunctionsParser
from langchain.prompts import SystemMessagePromptTemplate
from langchain.schema.output_parser import StrOutputParser
from permchain.channels import LastValue
from permchain.pregel import Pregel
# prompts
drafter_prompt = (
SystemMessagePromptTemplate.from_template(
"You are an expert on turtles, who likes to write in pirate-speak. You have been tasked by your editor with drafting a 100-word article answering the following question."
)
+ "Question:\n\n{question}"
)
reviser_prompt = (
SystemMessagePromptTemplate.from_template(
"You are an expert on turtles. You have been tasked by your editor with revising the following draft, which was written by a non-expert. You may follow the editor's notes or not, as you see fit."
)
+ "Draft:\n\n{draft}"
+ "Editor's notes:\n\n{notes}"
)
editor_prompt = (
SystemMessagePromptTemplate.from_template(
"You are an editor. You have been tasked with editing the following draft, which was written by a non-expert. Please accept the draft if it is good enough to publish, or send it for revision, along with your notes to guide the revision."
)
+ "Draft:\n\n{draft}"
)
editor_functions = [
{
"name": "revise",
"description": "Sends the draft for revision",
"parameters": {
"type": "object",
"properties": {
"notes": {
"type": "string",
"description": "The editor's notes to guide the revision.",
},
},
},
},
{
"name": "accept",
"description": "Accepts the draft",
"parameters": {
"type": "object",
"properties": {"ready": {"const": True}},
},
},
]
# llms
gpt3 = ChatOpenAI(model="gpt-3.5-turbo")
gpt4 = ChatOpenAI(model="gpt-4")
# chains
drafter_chain = drafter_prompt | gpt3 | StrOutputParser()
reviser_chain = reviser_prompt | gpt3 | StrOutputParser()
editor_chain = (
editor_prompt
| gpt4.bind(functions=editor_functions)
| JsonOutputFunctionsParser(args_only=False)
)
# state
question = LastValue[str]()
draft = LastValue[str]()
notes = LastValue[str]()
# application
drafter_node = Pregel.read(question=question) | drafter_chain | Pregel.write(draft)
reviser_node = (
Pregel.read(question=question, notes=notes, draft=draft)
| reviser_chain
| Pregel.write(draft)
)
editor_node = (
Pregel.read(draft=draft)
| editor_chain
| Pregel.write(
{notes: lambda x: x["arguments"]["notes"] if x["name"] == "revise" else None}
)
)
draft_revise_loop = Pregel(
(drafter_node, reviser_node, editor_node),
input=question,
output=draft,
)
# run
article = draft_revise_loop.invoke("What food do turtles eat?")
+86
View File
@@ -0,0 +1,86 @@
from abc import ABC, abstractmethod
from typing import Callable, FrozenSet, Generic, Self, Sequence, TypeVar
Value = TypeVar("Value")
Update = TypeVar("Update")
class EmptyChannelError(Exception):
pass
class InvalidUpdateError(Exception):
pass
class Channel(Generic[Value, Update], ABC):
def _empty(self) -> Self:
return self.__class__()
@abstractmethod
def _update(self, values: Sequence[Update]) -> None:
...
@abstractmethod
def _get(self) -> Value:
...
class BinaryOperatorAggregate(Generic[Value], Channel[Value, Value]):
def __init__(self, operator: Callable[[Value, Value], Value]):
self.operator = operator
def _empty(self) -> Self:
return self.__class__(self.operator)
def _update(self, values):
if not hasattr(self, "value"):
self.value = values[0]
values = values[1:]
for value in values:
self.value = self.operator(self.value, value)
def _get(self):
try:
return self.value
except AttributeError:
raise EmptyChannelError()
class LastValue(Generic[Value], Channel[Value, Value]):
def _update(self, values):
if len(values) != 1:
raise InvalidUpdateError()
self.value = values[-1]
def _get(self):
try:
return self.value
except AttributeError:
raise EmptyChannelError()
class Inbox(Generic[Value], Channel[Sequence[Value], Value]):
def _update(self, values):
self.queue = tuple(values)
def _get(self):
try:
return self.queue
except AttributeError:
raise EmptyChannelError()
class Set(Generic[Value], Channel[FrozenSet[Value], Value]):
def _update(self, values) -> None:
if not hasattr(self, "set"):
self.set = set()
self.set.update(values)
def _get(self) -> FrozenSet[Value]:
try:
return frozenset(self.set)
except AttributeError:
raise EmptyChannelError()
+453
View File
@@ -0,0 +1,453 @@
from __future__ import annotations
import asyncio
import concurrent.futures
import logging
from collections import defaultdict, deque
from typing import (
Any,
AsyncIterator,
Callable,
Generic,
Iterator,
Mapping,
Optional,
Sequence,
overload,
)
from langchain.callbacks.manager import (
CallbackManagerForChainRun,
AsyncCallbackManagerForChainRun,
)
from langchain.pydantic_v1 import Field
from langchain.schema.runnable import (
Runnable,
RunnableSerializable,
RunnableBinding,
RunnablePassthrough,
)
from langchain.schema.runnable.base import (
RunnableLike,
Other,
coerce_to_runnable,
RunnableEach,
)
from langchain.schema.runnable.config import (
RunnableConfig,
get_executor_for_config,
patch_config,
)
from langchain.schema.runnable.utils import ConfigurableFieldSpec, Input, Output
from permchain.channels import Channel, EmptyChannelError, Inbox
logger = logging.getLogger(__name__)
CONFIG_KEY_STEP = "__pregel_step"
CONFIG_KEY_WRITE = "__pregel_write"
class PregelInvoke(RunnableBinding):
channels: Mapping[str | None, Channel]
bound: Runnable[Any, Any] = Field(default_factory=RunnablePassthrough)
kwargs: Mapping[str, Any] = Field(default_factory=dict)
def __or__(
self,
other: Runnable[Any, Other]
| Callable[[Any], Other]
| Mapping[str, Runnable[Any, Other] | Callable[[Any], Other]],
) -> PregelInvoke:
if isinstance(self.bound, RunnablePassthrough):
return PregelInvoke(channels=self.channels, bound=coerce_to_runnable(other))
else:
return PregelInvoke(channels=self.channels, bound=self.bound | other)
def __ror__(
self,
other: Runnable[Other, Any]
| Callable[[Any], Other]
| Mapping[str, Runnable[Other, Any] | Callable[[Other], Any]],
) -> PregelInvoke:
raise NotImplementedError()
class PregelBatch(RunnableEach):
channel: Inbox
bound: Runnable[Any, Any] = Field(default_factory=RunnablePassthrough)
def __or__(
self,
other: Runnable[Any, Other]
| Callable[[Any], Other]
| Mapping[str, Runnable[Any, Other] | Callable[[Any], Other]],
) -> PregelBatch:
if isinstance(self.bound, RunnablePassthrough):
return PregelBatch(channel=self.channel, bound=coerce_to_runnable(other))
else:
return PregelBatch(channel=self.channel, bound=self.bound | other)
def __ror__(
self,
other: Runnable[Other, Any]
| Callable[[Any], Other]
| Mapping[str, Runnable[Other, Any] | Callable[[Other], Any]],
) -> PregelBatch:
raise NotImplementedError()
class PregelSink(RunnablePassthrough):
channels: Mapping[Channel, Runnable]
"""
Mapping of write channels to Runnables that return the value to be written,
or None to skip writing.
"""
max_steps: Optional[int]
@property
def config_specs(self) -> Sequence[ConfigurableFieldSpec]:
return [
ConfigurableFieldSpec(
id=CONFIG_KEY_STEP,
annotation=int,
),
ConfigurableFieldSpec(
id=CONFIG_KEY_WRITE,
annotation=Callable[[Sequence[tuple[Channel, Any]]], None],
),
]
def _write(self, input: Any, config: RunnableConfig) -> None:
step: int = config.get("configurable", {})[CONFIG_KEY_STEP]
if step >= self.max_steps:
return
write: Callable[[Sequence[tuple[Channel, Any]]], None] = config.get(
"configurable", {}
)[CONFIG_KEY_WRITE]
# TODO use runnable map to run this in parallel?
values = [(chan, r.invoke(input, config)) for chan, r in self.channels.items()]
values = [(chan, val) for chan, val in values if val is not None]
write(values)
return input
# TODO def _awrite()
class Pregel(Generic[Input, Output], RunnableSerializable[Input, Output]):
input: Channel[Any, Input]
output: Channel[Output, Any]
processes: Sequence[PregelInvoke | PregelBatch]
step_timeout: Optional[float] = None
def __init__(
self,
processes: Sequence[PregelInvoke | PregelBatch],
*,
input: Channel[Input, Any],
output: Channel[Output, Any],
step_timeout: Optional[float] = None,
**kwargs: Any,
):
super().__init__(
processes=processes,
input=input,
output=output,
step_timeout=step_timeout,
**kwargs,
)
@overload
@classmethod
def read(cls, __channel: Channel) -> PregelInvoke:
...
@overload
def read(cls, __channel: Mapping[str, Channel], **kwargs: Channel) -> PregelInvoke:
...
@classmethod
def read(
cls, __channel: Channel | Mapping[str, Channel], **kwargs: Channel
) -> PregelInvoke:
"""Runs process.invoke() each time channels are updated."""
return PregelInvoke(
channels=(
{None: __channel}
if isinstance(__channel, Channel)
else {**__channel, **kwargs}
)
)
@classmethod
def read_batch(cls, inbox: Inbox):
"""Runs process.batch() on the current contents of the inbox."""
return PregelBatch(channel=inbox)
@classmethod
def write(
cls,
channels: Channel | Mapping[Channel, RunnableLike],
*,
max_steps: Optional[int] = None,
):
return PregelSink(
channels=(
{channels: RunnablePassthrough()}
if isinstance(channels, Channel)
else {**channels}
),
max_steps=max_steps,
)
# TODO def write_each()
def _prepare_channels(self) -> Mapping[Channel, Channel]:
channels = {self.output: self.output._empty()}
for proc in self.processes:
if isinstance(proc, PregelInvoke):
for chan in proc.channels.values():
if chan not in channels:
channels[chan] = chan._empty()
elif isinstance(proc, PregelBatch):
if proc.channel not in channels:
channels[proc.channel] = proc.channel._empty()
else:
raise TypeError(
f"Received process {proc}, expected instance of PregelInvoke or PregelBatch"
)
if not channels:
return ValueError("Found 0 channels for Pregel run")
if self.input not in channels:
return ValueError("Input channel not being read from")
return channels
def _transform(
self,
input: Iterator[Input],
run_manager: CallbackManagerForChainRun,
config: RunnableConfig,
) -> Iterator[Output]:
processes = tuple(self.processes)
channels = self._prepare_channels()
next_tasks = _apply_writes_and_prepare_next_tasks(
processes, channels, deque((self.input, chunk) for chunk in input)
)
with get_executor_for_config(config) as executor:
# Similarly to Bulk Synchronous Parallel / Pregel model
# computation proceeds in steps, while there are channel updates
# channel updates from step N are only visible in step N+1
# channels are guaranteed to be immutable for the duration of the step,
# with channel updates applied only at the transition between steps
for step in range(config["recursion_limit"]):
# collect all writes to channels, without applying them yet
pending_writes = deque[tuple[Channel, Any]]()
# execute tasks, and wait for one to fail or all to finish
# each task is independent from all other concurrent tasks
done, inflight = concurrent.futures.wait(
(
executor.submit(
proc.invoke,
input,
patch_config(
config,
callbacks=run_manager.get_child(f"pregel:step:{step}"),
configurable={
# deque.extend is thread-safe
CONFIG_KEY_WRITE: pending_writes.extend,
CONFIG_KEY_STEP: step,
},
),
)
for proc, input in next_tasks
),
return_when=concurrent.futures.FIRST_EXCEPTION,
timeout=self.step_timeout,
)
while done:
# if any task failed
if exc := done.pop().exception():
# cancel all pending tasks
while inflight:
inflight.pop().cancel()
# raise the exception
raise exc
# TODO this is where retry of an entire step would happen
if inflight:
# if we got here means we timed out
while inflight:
# cancel all pending tasks
inflight.pop().cancel()
# raise timeout error
raise TimeoutError(f"Timed out at step {step}")
# apply writes to channels, decide on next step
next_tasks = _apply_writes_and_prepare_next_tasks(
processes, channels, pending_writes
)
# if any write to output channel in this step, yield current value
if any(chan is self.output for chan, _ in pending_writes):
yield channels[self.output]._get()
# if no more tasks, we're done
if not next_tasks:
break
# TODO clean up inflight futures if stream() is interrupted ?
# Test this first
# If this is needed implement with a weakset of futures and try/finally
async def _atransform(
self,
input: AsyncIterator[Input],
run_manager: AsyncCallbackManagerForChainRun,
config: RunnableConfig,
) -> AsyncIterator[Output]:
processes = tuple(self.processes)
channels = self._prepare_channels()
next_tasks = _apply_writes_and_prepare_next_tasks(
processes, channels, [(self.input, chunk) for chunk in input]
)
# Similarly to Bulk Synchronous Parallel / Pregel model
# computation proceeds in steps, while there are channel updates
# channel updates from step N are only visible in step N+1,
# channels are guaranteed to be immutable for the duration of the step,
# channel updates being applied only at the transition between steps
for step in range(config["recursion_limit"]):
# collect all writes to channels, without applying them yet
pending_writes = []
# execute tasks, and wait for one to fail or all to finish
# each task is independent from all other concurrent tasks
done, inflight = await asyncio.wait(
(
asyncio.create_task(
proc.ainvoke(
input,
patch_config(
config,
callbacks=run_manager.get_child(f"pregel:step:{step}"),
configurable={
# deque.extend is thread-safe
CONFIG_KEY_WRITE: pending_writes.extend,
CONFIG_KEY_STEP: step,
},
),
)
)
for proc, input in next_tasks
),
return_when=asyncio.FIRST_EXCEPTION,
timeout=self.step_timeout,
)
while done:
# if any task failed
if exc := done.pop().exception():
# cancel all pending tasks
while inflight:
inflight.pop().cancel()
# raise the exception
raise exc
# TODO this is where retry of an entire step would happen
if inflight:
# if we got here means we timed out
while inflight:
# cancel all pending tasks
inflight.pop().cancel()
# raise timeout error
raise TimeoutError(f"Timed out at step {step}")
# apply writes to channels, decide on next step
next_tasks = _apply_writes_and_prepare_next_tasks(
processes, channels, pending_writes
)
# if any write to output channel in this step, yield current value
if any(chan is self.output for chan, _ in pending_writes):
yield channels[self.output]._get()
# if no more tasks, we're done
if not next_tasks:
break
# TODO invoke() consumes stream() iterator and returns last value
# TODO ainvoke() consumes astream() iterator and returns last value
# TODO do we want api to subscribe to all channels?
def _apply_writes_and_prepare_next_tasks(
processes: Sequence[PregelInvoke | PregelBatch],
channels: Mapping[Channel, Channel],
pending_writes: Sequence[tuple[Channel, Any]],
) -> list[tuple[Runnable, Any]]:
pending_writes_by_channel: dict[Channel, list[Any]] = defaultdict(list)
# Group writes by channel
for chan, val in pending_writes:
pending_writes_by_channel[chan].append(val)
updated_channels: set[Channel] = set()
# Apply writes to channels
for chan, vals in pending_writes_by_channel.items():
if chan in channels:
channels[chan]._update(vals)
updated_channels.add(chan)
else:
logger.warning(f"Skipping write for channel {chan} which has no readers")
tasks: list[tuple[Runnable, Any]] = []
# Check if any processes should be run in next step
# If so, prepare the values to be passed to them
for proc in processes:
if isinstance(proc, PregelInvoke):
# If any of the channels read by this process were updated
if any(chan in updated_channels for chan in proc.channels.values()):
# If all channels read by this process have been initialized
try:
val = {
k: channels[chan]._get() for k, chan in proc.channels.items()
}
except EmptyChannelError:
continue
# Processes that subscribe to a single keyless channel get
# the value directly, instead of a dict
if list(proc.channels.keys()) == [None]:
tasks.append((proc, val[None]))
else:
tasks.append((proc, val))
elif isinstance(proc, PregelBatch):
# If the channel read by this process was updated
if proc.channel in updated_channels:
# Here we don't catch EmptyChannelError because the channel
# must be intialized if the previous `if` condition is true
val = channels[proc.channel]._get()
tasks.append((proc, val))
return tasks