Update for latest langchain, start to implement isolation and batch

This commit is contained in:
Nuno Campos
2023-08-18 12:30:55 +01:00
parent 7c867d8dac
commit e5532d59eb
5 changed files with 50 additions and 33 deletions
+1 -2
View File
@@ -1,8 +1,7 @@
from abc import ABC, abstractmethod
import asyncio
from abc import ABC, abstractmethod
from typing import Any, Callable
PubSubListener = Callable[[Any], None]
+1 -1
View File
@@ -1,5 +1,5 @@
from collections import defaultdict
import threading
from collections import defaultdict
from typing import Any
from permchain.connection import PubSubConnection, PubSubListener
+2
View File
@@ -0,0 +1,2 @@
CONFIG_GET_KEY = "pubsub_get"
CONFIG_SEND_KEY = "pubsub_send"
+41 -26
View File
@@ -1,8 +1,9 @@
from __future__ import annotations
from ctypes import Union
import queue
from abc import ABC
from concurrent.futures import CancelledError, Future, ThreadPoolExecutor
from concurrent.futures import CancelledError, Future
from functools import partial
from itertools import filterfalse
from typing import (
@@ -17,13 +18,17 @@ from typing import (
TypeVar,
)
from langchain.callbacks.manager import CallbackManager
from langchain.load.dump import dumpd
from langchain.load.serializable import Serializable
from langchain.schema.runnable import Runnable, RunnableConfig
from langchain.schema.runnable.base import patch_config
from langchain.schema.runnable import Runnable, RunnableConfig, patch_config
from langchain.schema.runnable.config import (
ensure_config,
get_callback_manager_for_config,
get_executor_for_config,
)
from permchain.connection import PubSubConnection
from permchain.constants import CONFIG_GET_KEY, CONFIG_SEND_KEY
from permchain.topic import INPUT_TOPIC, OUTPUT_TOPIC, RunnableSubscriber
T = TypeVar("T")
@@ -68,32 +73,34 @@ class PubSub(Serializable, Runnable[Any, Any], ABC):
collected.append(chunk)
return collected
def batch(
self,
input: Sequence[Any],
config: Optional[Union[RunnableConfig, List[RunnableConfig]]] = None,
) -> Any:
configs = self._get_config_list(config)
with get_executor_for_config(configs[0]) as executor:
return super().batch(
input,
[patch_config(config, executor=executor) for config in configs],
)
def stream(
self,
input: Any,
config: Optional[RunnableConfig] = None,
*,
max_concurrency: Optional[int] = None,
) -> Iterator[Any]:
input_processes, listener_processes = partition(
lambda r: r.topic.name == INPUT_TOPIC, self.processes
)
# setup callbacks
config = config or {}
callback_manager = CallbackManager.configure(
inheritable_callbacks=config.get("callbacks"),
local_callbacks=None,
verbose=False,
inheritable_tags=config.get("tags"),
local_tags=None,
inheritable_metadata=config.get("metadata"),
local_metadata=None,
)
config = ensure_config(config)
callback_manager = get_callback_manager_for_config(config)
# start the root run
run_manager = callback_manager.on_chain_start(dumpd(self), {"input": input})
with ThreadPoolExecutor(max_workers=max_concurrency) as executor:
with get_executor_for_config(config) as executor:
# Track inflight futures
inflight: Set[Future] = set()
# Track exceptions
@@ -101,16 +108,20 @@ class PubSub(Serializable, Runnable[Any, Any], ABC):
# Track output
output = IterableQueue()
# Namespace topics for each run
def prefix_topic_name(topic_name: str) -> str:
return f"{run_manager.run_id}/{topic_name}"
def send(topic_name: str, message: Any) -> None:
"""Send a message to a topic. Injected into config."""
if topic_name == OUTPUT_TOPIC:
output.put(message)
else:
self.connection.send(topic_name, message)
self.connection.send(prefix_topic_name(topic_name), message)
def cleanup_run(fut: Future) -> None:
"""Cleanup after a process runs."""
inflight.remove(fut)
inflight.discard(fut)
try:
exc = fut.exception()
@@ -146,10 +157,12 @@ class PubSub(Serializable, Runnable[Any, Any], ABC):
value,
config={
**patch_config(
config, run_manager.get_child(process.topic.name)
config,
callbacks=run_manager.get_child(process.topic.name),
executor=executor,
),
"send": send,
"get": get,
CONFIG_SEND_KEY: send,
CONFIG_GET_KEY: get,
},
)
@@ -159,7 +172,9 @@ class PubSub(Serializable, Runnable[Any, Any], ABC):
# Listen on all subscribed topics
for process in listener_processes:
self.connection.listen(process.topic.name, partial(run_once, process))
self.connection.listen(
prefix_topic_name(process.topic.name), partial(run_once, process)
)
# Run input processes once
for process in input_processes:
@@ -176,11 +191,11 @@ class PubSub(Serializable, Runnable[Any, Any], ABC):
final_output += chunk
finally:
# Cleanup
for fut in inflight:
fut.cancel()
while inflight:
inflight.pop().cancel()
for process in listener_processes:
self.connection.disconnect(process.topic.name)
self.connection.disconnect(prefix_topic_name(process.topic.name))
# Raise exceptions if any
if exceptions:
+5 -4
View File
@@ -1,7 +1,6 @@
from __future__ import annotations
from abc import ABC
from typing import (
Any,
Callable,
@@ -13,6 +12,7 @@ from typing import (
)
from langchain.load.serializable import Serializable
from langchain.pydantic_v1 import Field
from langchain.schema.runnable import (
Runnable,
RunnableBinding,
@@ -21,7 +21,8 @@ from langchain.schema.runnable import (
RunnableSequence,
)
from langchain.schema.runnable.base import Other, coerce_to_runnable
from pydantic import Field
from permchain.constants import CONFIG_GET_KEY, CONFIG_SEND_KEY
T = TypeVar("T")
T_in = TypeVar("T_in")
@@ -97,7 +98,7 @@ class RunnablePublisher(RunnablePassthrough[T]):
topic: Topic[T]
def invoke(self, input: T, config: Optional[RunnableConfigForPubSub] = None) -> T:
send = config.get("send", None)
send = config.get(CONFIG_SEND_KEY, None)
if send is not None:
send(self.topic.name, input)
return super().invoke(input, config)
@@ -117,7 +118,7 @@ class RunnableCurrentValue(Serializable, Runnable[Any, T]):
topic: Topic[T]
def invoke(self, input: T, config: Optional[RunnableConfigForPubSub] = None) -> T:
get = config.get("get", None)
get = config.get(CONFIG_GET_KEY, None)
if get is not None:
return get(self.topic.name)
else: