Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
14 changes: 14 additions & 0 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -389,6 +389,8 @@ struct Router {
assignment_mode: String,
max_payload_size: usize,
dp_aware: bool,
multimodal_tensor_transport: Option<String>,
multimodal_shm_min_bytes: Option<usize>,
dp_minimum_tokens_scheduler: bool,
api_key: Option<String>,
log_dir: Option<String>,
Expand Down Expand Up @@ -785,6 +787,12 @@ impl Router {
.maybe_storage_hook_wasm_path(self.storage_hook_wasm_path.as_deref())
.enable_wasm(self.enable_wasm)
.dp_aware(self.dp_aware)
.multimodal_tensor_transport(
self.multimodal_tensor_transport
.as_deref()
.and_then(config::TransportMode::parse),
)
Comment thread
slin1237 marked this conversation as resolved.
Outdated
Comment thread
slin1237 marked this conversation as resolved.
Outdated
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
.multimodal_shm_min_bytes(self.multimodal_shm_min_bytes)
.routing_key_override(config::RoutingKeyOverrideConfig {
enabled: self.routing_key_override,
eviction_interval_secs: self.eviction_interval_secs,
Expand Down Expand Up @@ -834,6 +842,8 @@ impl Router {
assignment_mode = String::from("random"),
max_payload_size = 512 * 1024 * 1024,
dp_aware = false,
multimodal_tensor_transport = None,
multimodal_shm_min_bytes = None,
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
dp_minimum_tokens_scheduler = false,
api_key = None,
log_dir = None,
Expand Down Expand Up @@ -959,6 +969,8 @@ impl Router {
assignment_mode: String,
max_payload_size: usize,
dp_aware: bool,
multimodal_tensor_transport: Option<String>,
multimodal_shm_min_bytes: Option<usize>,
dp_minimum_tokens_scheduler: bool,
api_key: Option<String>,
log_dir: Option<String>,
Expand Down Expand Up @@ -1100,6 +1112,8 @@ impl Router {
assignment_mode,
max_payload_size,
dp_aware,
multimodal_tensor_transport,
multimodal_shm_min_bytes,
dp_minimum_tokens_scheduler,
api_key,
log_dir,
Expand Down
17 changes: 17 additions & 0 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,8 @@ class RouterArgs:
max_payload_size: int = 512 * 1024 * 1024 # 512MB default for large batches
bucket_adjust_interval_secs: int = 5
dp_aware: bool = False
multimodal_tensor_transport: str | None = None
multimodal_shm_min_bytes: int | None = None
routing_key_override: bool = False
dp_minimum_tokens_scheduler: bool = False
enable_igw: bool = False # Enable IGW (Inter-Gateway) mode for multi-model support
Expand Down Expand Up @@ -565,6 +567,21 @@ def add_cli_args(
help="Interval in seconds between load monitor checks for PowerOfTwo routing (default: 10)",
)

# Multimodal tensor transport
parser.add_argument(
f"--{prefix}multimodal-tensor-transport",
type=str,
choices=["inline", "shm", "auto"],
default=RouterArgs.multimodal_tensor_transport,
help="Multimodal tensor transport mode: inline (default), shm, or auto",
)
parser.add_argument(
f"--{prefix}multimodal-shm-min-bytes",
type=int,
default=RouterArgs.multimodal_shm_min_bytes,
help="Minimum multimodal tensor size (bytes) before the SHM transport is used",
)

# Logging configuration
logging_group.add_argument(
f"--{prefix}log-dir",
Expand Down
66 changes: 66 additions & 0 deletions crates/protocols/src/worker.rs
Original file line number Diff line number Diff line change
Expand Up @@ -670,6 +670,17 @@ pub struct WorkerSpec {
/// Falls back to the global `load_monitor_interval_secs` from router config.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub load_monitor_interval_secs: Option<u64>,

/// Per-worker multimodal tensor transport override (`inline` | `shm` | `auto`).
/// Overrides the router-level `multimodal_tensor_transport` for this worker
/// (e.g. force `shm` for a co-located worker, `inline` for a remote one).
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multimodal_tensor_transport: Option<TransportMode>,

/// Per-worker minimum multimodal tensor size (bytes) before the SHM transport
/// is used. Overrides the router-level `multimodal_shm_min_bytes`.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multimodal_shm_min_bytes: Option<usize>,
}

impl WorkerSpec {
Expand Down Expand Up @@ -700,10 +711,65 @@ impl WorkerSpec {
resilience: ResilienceUpdate::default(),
max_connection_attempts: default_max_connection_attempts(),
load_monitor_interval_secs: None,
multimodal_tensor_transport: None,
multimodal_shm_min_bytes: None,
}
}
}

/// Multimodal tensor transport mode for large payloads.
///
/// - `Inline`: always carry tensor bytes in the gRPC message.
/// - `Shm`: use same-host `/dev/shm` when SMG can write it.
/// - `Auto`: use `/dev/shm` only when the receiving worker is verified to share
/// SMG's `/dev/shm`; otherwise fall back to inline.
#[derive(
Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize, schemars::JsonSchema,
)]
#[serde(rename_all = "lowercase")]
pub enum TransportMode {
#[default]
Inline,
Shm,
Auto,
}

impl TransportMode {
/// Parse from a case-insensitive string (`inline` | `shm` | `auto`).
pub fn parse(value: &str) -> Option<Self> {
match value.trim().to_ascii_lowercase().as_str() {
"inline" => Some(Self::Inline),
"shm" => Some(Self::Shm),
"auto" => Some(Self::Auto),
_ => None,
}
}

/// Canonical lowercase name.
pub fn as_str(self) -> &'static str {
match self {
Self::Inline => "inline",
Self::Shm => "shm",
Self::Auto => "auto",
}
}
}

impl std::fmt::Display for TransportMode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}

impl std::str::FromStr for TransportMode {
type Err = String;

fn from_str(s: &str) -> Result<Self, Self::Err> {
Self::parse(s)
.ok_or_else(|| format!("invalid transport mode '{s}'; expected inline|shm|auto"))
}
}

// ── API types ───────────────────────────────────────────────────────

/// Worker information for API responses.
Expand Down
22 changes: 17 additions & 5 deletions docs/reference/configuration.md
Original file line number Diff line number Diff line change
Expand Up @@ -937,16 +937,28 @@ smg \
| `JWT_JWKS_URI` | `--jwt-jwks-uri` | JWKS URI |
| `CONTROL_PLANE_API_KEYS` | `--control-plane-api-keys` | Control plane API keys |

### TokenSpeed Multimodal Tensor Transport
### Multimodal Tensor Transport

These env-only variables tune how the router ships preprocessed multimodal
tensors (image/video encoder inputs) to a TokenSpeed worker. They do not affect
Controls how the router ships preprocessed multimodal tensors (image/video
encoder inputs and model-specific tensors) to the worker. This does not affect
accuracy — the inline and shared-memory paths produce byte-identical tensors.
Shared memory is currently used by the TokenSpeed backend.

Resolution precedence (highest first): per-worker `WorkerSpec` override →
router config / CLI flag → `SMG_MM_*` environment variable → built-in default.

| CLI Flag | Config / WorkerSpec field | Default | Description |
|----------|---------------------------|---------|-------------|
| `--multimodal-tensor-transport` | `multimodal_tensor_transport` | `inline` | Transport for large MM tensors: `inline` (gRPC bytes), `shm` (use `/dev/shm` whenever the router can write it — the operator asserts co-location), or `auto` (use `/dev/shm` only when the worker is *verified* to share it). In `auto`, the router compares the worker's advertised `/dev/shm` namespace token (`GetServerInfo`) to its own and uses SHM only on a match; otherwise it falls back to inline. |
| `--multimodal-shm-min-bytes` | `multimodal_shm_min_bytes` | `65536` | Minimum tensor size (bytes) before the SHM path is used; smaller tensors stay inline. |

Both settings can be overridden per worker via the matching `WorkerSpec` fields
(e.g. force `shm` for a co-located worker and `inline` for a remote one).

| Environment Variable | Default | Description |
|---------------------|---------|-------------|
| `SMG_TOKENSPEED_MM_TENSOR_TRANSPORT` | `inline` | Transport for large MM tensors: `inline` (gRPC bytes), `shm` (always use `/dev/shm`), or `auto` (use `/dev/shm` only when the worker is *verified* to share it). In `auto`, the router compares the worker's advertised `/dev/shm` namespace token (`GetServerInfo`) to its own and uses SHM only on a match; otherwise it falls back to inline. No locality configuration is needed. |
| `SMG_TOKENSPEED_MM_SHM_MIN_BYTES` | `65536` | Minimum tensor size (bytes) before the SHM path is used; smaller tensors stay inline. |
| `SMG_MM_TENSOR_TRANSPORT` | `inline` | Env fallback for `--multimodal-tensor-transport`. Legacy alias: `SMG_TOKENSPEED_MM_TENSOR_TRANSPORT`. |
| `SMG_MM_SHM_MIN_BYTES` | `65536` | Env fallback for `--multimodal-shm-min-bytes`. Legacy alias: `SMG_TOKENSPEED_MM_SHM_MIN_BYTES`. |
| `SMG_LOG_MM_TIMING` | `false` | Log per-stage multimodal preprocessing/assembly timing at `INFO`. Accepts `1`/`true`/`yes`. |

The TokenSpeed gRPC servicer (worker side) reads two companion variables:
Expand Down
13 changes: 13 additions & 0 deletions model_gateway/src/config/builder.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
use std::collections::HashMap;

use openai_protocol::worker::TransportMode;
use smg_mcp::McpConfig;

use super::{
Expand Down Expand Up @@ -213,6 +214,18 @@ impl RouterConfigBuilder {
self
}

/// Global multimodal tensor transport mode (per-worker specs can override).
pub fn multimodal_tensor_transport(mut self, mode: Option<TransportMode>) -> Self {
self.config.multimodal_tensor_transport = mode;
self
}

/// Global minimum multimodal tensor size (bytes) before SHM transport is used.
pub fn multimodal_shm_min_bytes(mut self, bytes: Option<usize>) -> Self {
self.config.multimodal_shm_min_bytes = bytes;
self
}

// ==================== Rate Limiting ====================

pub fn max_concurrent_requests(mut self, max: i32) -> Self {
Expand Down
13 changes: 13 additions & 0 deletions model_gateway/src/config/types.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
use std::collections::HashMap;

use openai_protocol::worker::HealthCheckConfig as ProtocolHealthCheckConfig;
pub use openai_protocol::worker::TransportMode;
use serde::{Deserialize, Serialize};
// Re-export storage config types from data_connector
pub use smg_data_connector::{
Expand Down Expand Up @@ -43,6 +44,16 @@ pub struct RouterConfig {
/// observability from routing.
#[serde(default)]
pub engine_metrics: bool,
/// Global multimodal tensor transport mode (`inline` | `shm` | `auto`).
/// Per-worker `WorkerSpec.multimodal_tensor_transport` overrides this; when
/// unset, falls back to `SMG_MM_TENSOR_TRANSPORT`, then `inline`.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multimodal_tensor_transport: Option<TransportMode>,
/// Global minimum multimodal tensor size (bytes) before SHM transport is used.
/// Per-worker `WorkerSpec.multimodal_shm_min_bytes` overrides this; falls back
/// to `SMG_MM_SHM_MIN_BYTES`, then 64 KiB.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub multimodal_shm_min_bytes: Option<usize>,
pub dp_aware: bool,
#[serde(default)]
pub dp_minimum_tokens_scheduler: bool,
Expand Down Expand Up @@ -743,6 +754,8 @@ impl Default for RouterConfig {
worker_startup_check_interval_secs: 30,
load_monitor_interval_secs: 10,
engine_metrics: false,
multimodal_tensor_transport: None,
multimodal_shm_min_bytes: None,
dp_aware: false,
dp_minimum_tokens_scheduler: false,
api_key: None,
Expand Down
51 changes: 51 additions & 0 deletions model_gateway/src/main.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
use std::collections::HashMap;

use clap::{ArgAction, Parser, Subcommand, ValueEnum};
use openai_protocol::worker::TransportMode;
use rand::{distr::Alphanumeric, RngExt};
use smg::{
config::{
Expand Down Expand Up @@ -146,6 +147,12 @@ enum Commands {
},
}

/// Parse the `--multimodal-tensor-transport` value into a `TransportMode`.
fn parse_transport_mode(value: &str) -> Result<TransportMode, String> {
TransportMode::parse(value)
.ok_or_else(|| format!("invalid value '{value}'; expected inline, shm, or auto"))
}

#[derive(Parser, Debug)]
struct CliArgs {
// ==================== Worker Configuration ====================
Expand Down Expand Up @@ -308,6 +315,17 @@ struct CliArgs {
#[arg(long, default_value_t = false, help_heading = "Load Monitoring")]
engine_metrics: bool,

/// Multimodal tensor transport mode: `inline` (default), `shm` (same-host
/// /dev/shm), or `auto` (shm only when the worker shares /dev/shm). A
/// per-worker `WorkerSpec.multimodal_tensor_transport` overrides this.
#[arg(long, value_parser = parse_transport_mode, help_heading = "Multimodal")]
multimodal_tensor_transport: Option<TransportMode>,

/// Minimum multimodal tensor size (bytes) before the SHM transport is used.
/// Overridable per worker via `WorkerSpec.multimodal_shm_min_bytes`.
#[arg(long, help_heading = "Multimodal")]
multimodal_shm_min_bytes: Option<usize>,

// ==================== Service Discovery (Kubernetes) ====================
/// Enable Kubernetes service discovery
#[arg(
Expand Down Expand Up @@ -1338,6 +1356,8 @@ impl CliArgs {
.worker_startup_check_interval_secs(self.worker_startup_check_interval)
.load_monitor_interval_secs(self.load_monitor_interval)
.engine_metrics(self.engine_metrics)
.multimodal_tensor_transport(self.multimodal_tensor_transport)
.multimodal_shm_min_bytes(self.multimodal_shm_min_bytes)
.max_concurrent_requests(self.max_concurrent_requests)
.queue_size(self.queue_size)
.queue_timeout_secs(self.queue_timeout_secs)
Expand Down Expand Up @@ -1708,6 +1728,37 @@ mod tests {
);
}

/// The multimodal transport flags must reach both `RouterConfig` and the
/// wrapped `ServerConfig.router_config`. Two-path config-plumbing guard.
#[test]
fn multimodal_transport_flows_into_both_configs() {
let cli = cli_args_from(&[
"--multimodal-tensor-transport",
"shm",
"--multimodal-shm-min-bytes",
"1024",
]);

let router_config = cli.to_router_config(vec![], vec![]).unwrap();
assert_eq!(
router_config.multimodal_tensor_transport,
Some(TransportMode::Shm),
"transport mode must reach RouterConfig via to_router_config"
);
assert_eq!(router_config.multimodal_shm_min_bytes, Some(1024));

let server_config = cli.to_server_config(router_config).unwrap();
assert_eq!(
server_config.router_config.multimodal_tensor_transport,
Some(TransportMode::Shm),
"transport mode must survive into ServerConfig via to_server_config"
);
assert_eq!(
server_config.router_config.multimodal_shm_min_bytes,
Some(1024)
);
}

/// Default is off: the flag stays false through both conversions so
/// existing deployments keep the routing-gated polling behavior.
#[test]
Expand Down
9 changes: 7 additions & 2 deletions model_gateway/src/routers/grpc/epd_encode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -77,15 +77,17 @@ pub(crate) enum PreparedEncodeItem {
TokenSpeed {
item: Option<TokenSpeedMultimodalItem>,
shm_enabled: bool,
shm_min_bytes: usize,
cleanup_on_drop: bool,
},
}

impl PreparedEncodeItem {
fn tokenspeed(item: TokenSpeedMultimodalItem, shm_enabled: bool) -> Self {
fn tokenspeed(item: TokenSpeedMultimodalItem, shm_enabled: bool, shm_min_bytes: usize) -> Self {
Self::TokenSpeed {
item: Some(item),
shm_enabled,
shm_min_bytes,
cleanup_on_drop: true,
}
}
Expand All @@ -99,6 +101,7 @@ impl PreparedEncodeItem {
Self::TokenSpeed {
item,
shm_enabled,
shm_min_bytes,
cleanup_on_drop,
} => {
let item = item
Expand All @@ -111,6 +114,7 @@ impl PreparedEncodeItem {
TokenSpeedMultimodalData {
items: vec![item],
shm_enabled: *shm_enabled,
shm_min_bytes: *shm_min_bytes,
}
.into_proto(),
),
Expand Down Expand Up @@ -254,10 +258,11 @@ fn prepare_tokenspeed_items(
) -> Result<Vec<PreparedEncodeItem>> {
let tokenspeed_mm = assemble_tokenspeed(precomputed, workers, false)?;
let shm_enabled = tokenspeed_mm.shm_enabled;
let shm_min_bytes = tokenspeed_mm.shm_min_bytes;
Ok(tokenspeed_mm
.items
.into_iter()
.map(|item| PreparedEncodeItem::tokenspeed(item, shm_enabled))
.map(|item| PreparedEncodeItem::tokenspeed(item, shm_enabled, shm_min_bytes))
.collect())
}

Expand Down
Loading
Loading