Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 86 additions & 0 deletions examples/base/local/deployment_config_event_driven.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
# Event-driven deployment example
# Instead of polling at a fixed rate, the pipeline triggers automatically
# whenever a monitored PV value changes.
#
# Usage:
# pl run --config ./examples/base/local/deployment_config_event_driven.yaml
#
# Test by writing to a PV:
# pvput ML:LOCAL:TEST_A 3.14
# pvput ML:LOCAL:TEST_B 2.71
#
# The model will only evaluate when an input PV changes, rather than
# every N seconds.

deployment:
type: "event_driven"
min_monitor_interval: 0.01 # throttle: at most one update per PV every 100 ms
on_change_only: true # skip updates where the value hasn't changed

modules:
p4p_server:
name: "p4p_server"
type: "interface.p4p_server"
pub:
- "in_interface"
sub:
- "get_all"
- "out_transformer"
module_args: None
config:
EPICS_PVA_NAME_SERVERS: "localhost:5075"
variables:
ML:LOCAL:TEST_B:
proto: pva
name: ML:LOCAL:TEST_B
ML:LOCAL:TEST_A:
proto: pva
name: ML:LOCAL:TEST_A
ML:LOCAL:TEST_S:
proto: pva
name: ML:LOCAL:TEST_S

input_transformer:
name: "input_transformer"
type: "transformer.SimpleTransformer"
pub: "in_transformer"
sub: "in_interface"
module_args: None
config:
symbols:
- "ML:LOCAL:TEST_B"
- "ML:LOCAL:TEST_A"
variables:
x:
formula: "ML:LOCAL:TEST_A * 2 + 10"
y:
formula: "ML:LOCAL:TEST_B + 120"

model:
name: "model"
type: "model.SimpleModel"
pub: "model"
sub: "in_transformer"
module_args: None
config:
type: "LocalModelGetter"
args:
model_path: "examples/base/local/model_definition_event_driven.py"
model_factory_class: "ModelFactory"
variables:
max:
type: "scalar"

output_transformer:
name: "output_transformer"
type: "transformer.SimpleTransformer"
pub: "out_transformer"
sub: "model"
module_args:
unpack_data: True
config:
symbols:
- "output"
variables:
ML:LOCAL:TEST_S:
formula: "output"
130 changes: 130 additions & 0 deletions examples/base/local/deployment_config_multihead_event_driven.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
# Multi-headed event-driven deployment example
# Two input interfaces (FastAPI + p4p server) feed the same pipeline.
# The model evaluates whenever EITHER interface receives new values.
#
# Usage:
# pl run --config ./examples/base/local/deployment_config_multihead_event_driven.yaml -p
#
# Test via PVAccess:
# spput ML:LOCAL:TEST_A 3.14
# spput ML:LOCAL:TEST_B 2.71
#
# Test via HTTP:
# curl -X POST http://localhost:8001/submit -H 'Content-Type: application/json' -d '{"variables":{"ML:LOCAL:TEST_A":{"value":3.14},"ML:LOCAL:TEST_B":{"value":2.71}}}'
#
# Monitor output PV:
# spmonitor ML:LOCAL:TEST_S

deployment:
type: "event_driven"
min_monitor_interval: 0.01 # throttle per PV (seconds)
on_change_only: false

modules:
# ── PVAccess server interface ────────────────────────────────────────
p4p_server:
name: "p4p_server"
type: "interface.p4p_server"
pub:
- "in_interface_1"
sub:
- "get_all"
- "out_transformer"
- "in_interface_0"
module_args: None
config:
EPICS_PVA_NAME_SERVERS: "localhost:5075"
variables:
ML:LOCAL:TEST_A:
proto: pva
name: ML:LOCAL:TEST_A
ML:LOCAL:TEST_B:
proto: pva
name: ML:LOCAL:TEST_B
ML:LOCAL:TEST_S:
proto: pva
name: ML:LOCAL:TEST_S

# ── FastAPI HTTP interface ───────────────────────────────────────────
fastapi_server:
name: "fastapi_server"
type: "interface.fastapi_server"
pub:
- "in_interface_0"
sub:
- "get_all"
- "out_transformer"
- "in_interface_1"
module_args: None
config:
name: "fastapi_server"
host: "0.0.0.0"
port: 8001
start_server: true
wait_for_server_start: true
startup_timeout_s: 5.0
input_queue_max: 1000
output_queue_max: 1000
variables:
ML:LOCAL:TEST_A:
mode: in
type: scalar
default: 0.0
ML:LOCAL:TEST_B:
mode: in
type: scalar
default: 0.0
ML:LOCAL:TEST_S:
mode: out
type: scalar
default: 0.0

# ── Input transformer ───────────────────────────────────────────────
input_transformer:
name: "input_transformer"
type: "transformer.SimpleTransformer"
pub: "in_transformer"
sub:
- "in_interface_0"
- "in_interface_1"
module_args: None
config:
symbols:
- "ML:LOCAL:TEST_A"
- "ML:LOCAL:TEST_B"
variables:
x:
formula: "ML:LOCAL:TEST_A"
y:
formula: "ML:LOCAL:TEST_B"

# ── Model ────────────────────────────────────────────────────────────
model:
name: "model"
type: "model.SimpleModel"
pub: "model"
sub: "in_transformer"
module_args: None
config:
type: "LocalModelGetter"
args:
model_path: "examples/base/local/model_definition_event_driven.py"
model_factory_class: "ModelFactory"
variables:
max:
type: "scalar"

# ── Output transformer ──────────────────────────────────────────────
output_transformer:
name: "output_transformer"
type: "transformer.SimpleTransformer"
pub: "out_transformer"
sub: "model"
module_args:
unpack_data: True
config:
symbols:
- "output"
variables:
ML:LOCAL:TEST_S:
formula: "output"
51 changes: 51 additions & 0 deletions examples/base/local/model_definition_event_driven.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
import torch
import os
import time
import logging

logger = logging.getLogger(__name__)


class ModelFactory:
def __init__(self):
os.environ['PYTHONPATH'] = os.path.abspath(
os.path.join(os.path.dirname(__file__), '..', '..', '..')
)
self.model = SimpleModel()
model_path = 'examples/base/local/model.pth'
if os.path.exists(model_path):
self.model.load_state_dict(torch.load(model_path))
logger.info('Model loaded successfully.')
else:
logger.warning(
f"Model file '{model_path}' not found. Using untrained model."
)
logger.info('ModelFactory initialized (event-driven mode)')

def get_model(self):
return self.model


class SimpleModel(torch.nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.linear1 = torch.nn.Linear(2, 10)
self.linear2 = torch.nn.Linear(10, 1)
self._eval_count = 0

def forward(self, x):
x = torch.relu(self.linear1(x))
x = self.linear2(x)
return x

def evaluate(self, x: dict) -> dict:
self._eval_count += 1
logger.info(
f'[event-driven] evaluate #{self._eval_count} triggered at '
f'{time.strftime("%H:%M:%S")} | inputs: x={x.get("x")}, y={x.get("y")}'
)
input_tensor = torch.tensor([x['x'], x['y']], dtype=torch.float32)
output_tensor = self.forward(input_tensor)
result = output_tensor.item()
logger.info(f'[event-driven] output: {result}')
return {'output': result}
47 changes: 44 additions & 3 deletions poly_lithic/src/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,14 +71,15 @@ def load_env_config(env_path):
raise e


async def model_main(args, config, broker):
async def model_main(args, config, broker, builder=None):
"""
Main async function for running the model manager.

Args:
args: Parsed arguments namespace
config: Configuration object
broker: Broker instance
builder: Builder instance (needed for event_driven mode)
"""
logger = get_logger()
logger.info('Starting model manager')
Expand Down Expand Up @@ -106,6 +107,37 @@ async def model_main(args, config, broker):

await asyncio.sleep(0.01)

elif config.deployment.type == 'event_driven':
# Start PV monitors on all input interfaces
input_observers = builder.get_input_interface_observers()
if not input_observers:
raise ValueError('No input InterfaceObservers found for event-driven mode')

# Seed all transformer inputs with current PV values so the
# all-inputs-present gate passes on the first monitor event.
broker.get_all()
while broker.queue:
broker.parse_queue()
logger.info('Seeded transformer inputs via get_all')

for obs in input_observers:
obs.start_monitors(
broker,
min_interval=config.deployment.min_monitor_interval,
on_change_only=config.deployment.on_change_only,
)
logger.info(f'Started monitors on {len(input_observers)} input interface(s)')

while True:
if len(broker.queue) > 0:
broker.parse_queue()

if args.one_shot:
logger.info('One shot mode, exiting')
break

await asyncio.sleep(0.01)

else:
raise Exception(f'Deployment type "{config.deployment.type}" not supported')

Expand Down Expand Up @@ -218,10 +250,19 @@ def run_model(config, model_getter, debug, env, one_shot, publish, requirements)

# Import heavy dependencies only when needed
from poly_lithic.src.utils.builder import Builder
from poly_lithic.src.utils.trace_store import TraceStore
from poly_lithic.src.utils.trace_server import start_trace_server

click.echo('Building model manager...')
builder = Builder(config)
broker = builder.build()

trace_store = TraceStore(maxlen=builder.config.deployment.trace_buffer_size)
broker = builder.build(trace_store=trace_store)

# Start tracing API server
trace_port = int(os.environ.get('TRACE_PORT', builder.config.deployment.trace_port))
start_trace_server(trace_store, port=trace_port)
logger.info(f'Tracing API server started on port {trace_port}')

if requirements:
click.echo('Requirements-only mode - exiting after installation')
Expand All @@ -240,7 +281,7 @@ def run_model(config, model_getter, debug, env, one_shot, publish, requirements)
)

logger.info('Starting model manager main loop')
asyncio.run(model_main(args, builder.config, broker))
asyncio.run(model_main(args, builder.config, broker, builder=builder))

except KeyboardInterrupt:
click.echo('\n\nInterrupted by user')
Expand Down
11 changes: 11 additions & 0 deletions poly_lithic/src/config/config_object.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,17 @@ def validate_module_args(cls, v):
class DeploymentConfig(pydantic.BaseModel):
type: str
rate: Optional[Union[float, int]] = None
min_monitor_interval: float = 0.0
on_change_only: bool = False
trace_buffer_size: int = 10000
trace_port: int = 8100

@pydantic.field_validator('type')
def validate_type(cls, v):
allowed = {'continuous', 'event_driven'}
if v not in allowed:
raise ValueError(f'deployment type must be one of {allowed}, got {v!r}')
return v


class ConfigObject(pydantic.BaseModel):
Expand Down
2 changes: 2 additions & 0 deletions poly_lithic/src/interfaces/BaseInterface.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@ def save(self, data, **kwargs):


class BaseInterface(ABC):
supports_monitor: bool = False

@abstractmethod
def __init__(self, config):
pass
Expand Down
Loading
Loading