parallel
This commit is contained in:
@@ -4,10 +4,10 @@ import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
from minillmflow import AsyncNode, BatchAsyncFlow
|
||||
from minillmflow import AsyncNode, AsyncBatchFlow
|
||||
|
||||
class AsyncDataProcessNode(AsyncNode):
|
||||
def prep(self, shared_storage):
|
||||
async def prep_async(self, shared_storage):
|
||||
key = self.params.get('key')
|
||||
data = shared_storage['input_data'][key]
|
||||
if 'results' not in shared_storage:
|
||||
@@ -34,8 +34,8 @@ class TestAsyncBatchFlow(unittest.TestCase):
|
||||
|
||||
def test_basic_async_batch_processing(self):
|
||||
"""Test basic async batch processing with multiple keys"""
|
||||
class SimpleTestAsyncBatchFlow(BatchAsyncFlow):
|
||||
def prep(self, shared_storage):
|
||||
class SimpleTestAsyncBatchFlow(AsyncBatchFlow):
|
||||
async def prep_async(self, shared_storage):
|
||||
return [{'key': k} for k in shared_storage['input_data'].keys()]
|
||||
|
||||
shared_storage = {
|
||||
@@ -58,8 +58,8 @@ class TestAsyncBatchFlow(unittest.TestCase):
|
||||
|
||||
def test_empty_async_batch(self):
|
||||
"""Test async batch processing with empty input"""
|
||||
class EmptyTestAsyncBatchFlow(BatchAsyncFlow):
|
||||
def prep(self, shared_storage):
|
||||
class EmptyTestAsyncBatchFlow(AsyncBatchFlow):
|
||||
async def prep_async(self, shared_storage):
|
||||
return [{'key': k} for k in shared_storage['input_data'].keys()]
|
||||
|
||||
shared_storage = {
|
||||
@@ -73,8 +73,8 @@ class TestAsyncBatchFlow(unittest.TestCase):
|
||||
|
||||
def test_async_error_handling(self):
|
||||
"""Test error handling during async batch processing"""
|
||||
class ErrorTestAsyncBatchFlow(BatchAsyncFlow):
|
||||
def prep(self, shared_storage):
|
||||
class ErrorTestAsyncBatchFlow(AsyncBatchFlow):
|
||||
async def prep_async(self, shared_storage):
|
||||
return [{'key': k} for k in shared_storage['input_data'].keys()]
|
||||
|
||||
shared_storage = {
|
||||
@@ -110,8 +110,8 @@ class TestAsyncBatchFlow(unittest.TestCase):
|
||||
await asyncio.sleep(0.01)
|
||||
return "done"
|
||||
|
||||
class NestedAsyncBatchFlow(BatchAsyncFlow):
|
||||
def prep(self, shared_storage):
|
||||
class NestedAsyncBatchFlow(AsyncBatchFlow):
|
||||
async def prep_async(self, shared_storage):
|
||||
return [{'key': k} for k in shared_storage['input_data'].keys()]
|
||||
|
||||
# Create inner flow
|
||||
@@ -147,8 +147,8 @@ class TestAsyncBatchFlow(unittest.TestCase):
|
||||
shared_storage['results'][key] = shared_storage['input_data'][key] * multiplier
|
||||
return "done"
|
||||
|
||||
class CustomParamAsyncBatchFlow(BatchAsyncFlow):
|
||||
def prep(self, shared_storage):
|
||||
class CustomParamAsyncBatchFlow(AsyncBatchFlow):
|
||||
async def prep_async(self, shared_storage):
|
||||
return [{
|
||||
'key': k,
|
||||
'multiplier': i + 1
|
||||
|
||||
@@ -17,7 +17,7 @@ class AsyncNumberNode(AsyncNode):
|
||||
super().__init__()
|
||||
self.number = number
|
||||
|
||||
def prep(self, shared_storage):
|
||||
async def prep_async(self, shared_storage):
|
||||
# Synchronous work is allowed inside an AsyncNode,
|
||||
# but final 'condition' is determined by post_async().
|
||||
shared_storage['current'] = self.number
|
||||
@@ -34,7 +34,7 @@ class AsyncIncrementNode(AsyncNode):
|
||||
"""
|
||||
Demonstrates incrementing the 'current' value asynchronously.
|
||||
"""
|
||||
def prep(self, shared_storage):
|
||||
async def prep_async(self, shared_storage):
|
||||
shared_storage['current'] = shared_storage.get('current', 0) + 1
|
||||
return "incremented"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user