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
62 changes: 53 additions & 9 deletions mplang/backends/spu_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@

import logging
import os
from dataclasses import dataclass, field
from dataclasses import dataclass, field, replace
from typing import TYPE_CHECKING, Any

import spu.api as spu_api
Expand Down Expand Up @@ -117,19 +117,53 @@ def __init__(
# Optional shared infrastructure (for per-request isolation via link.spawn)
self._infra = infra
self._brpc_config = brpc_config or BrpcLinkConfig()
# Key: (local_rank, world_size, protocol, field, link_mode, spu_endpoints, brpc_config)
# ``brpc_config`` participates only when link_mode == "brpc" (else None) so that
# callers passing custom configs don't silently reuse the first config's link.
# Key: (local_rank, world_size, protocol, field, fxp_fraction_bits,
# link_mode, spu_endpoints, brpc_config). ``brpc_config`` participates
# only when link_mode == "brpc" (else None) so callers passing custom
# configs don't silently reuse the first config's link.
# Value: (Runtime, Io)
self._runtimes: dict[
tuple[
int, int, str, str, str, tuple[str, ...] | None, BrpcLinkConfig | None
int,
int,
str,
str,
int,
str,
tuple[str, ...] | None,
BrpcLinkConfig | None,
],
tuple[spu_api.Runtime, spu_api.Io],
] = {}
# Local template link cache (used when no WorkerInfra is provided)
self._template_links: dict[tuple, libspu.link.Context] = {}

def _effective_brpc_config(self, config: spu.SPUConfig) -> BrpcLinkConfig:
"""Merge per-SPU link_desc over this state's default brpc config."""
link_desc = config.link_desc
if link_desc is None:
return self._brpc_config

overrides: dict[str, Any] = {}
if link_desc.brpc_channel_protocol is not None:
overrides["protocol"] = link_desc.brpc_channel_protocol
if link_desc.brpc_channel_connection_type is not None:
overrides["connection_type"] = link_desc.brpc_channel_connection_type
if link_desc.recv_timeout_ms is not None:
overrides["recv_timeout_ms"] = link_desc.recv_timeout_ms
if link_desc.http_max_payload_size is not None:
overrides["http_max_payload_size"] = link_desc.http_max_payload_size
if link_desc.http_timeout_ms is not None:
overrides["http_timeout_ms"] = link_desc.http_timeout_ms
if link_desc.connect_retry_times is not None:
overrides["connect_retry_times"] = link_desc.connect_retry_times
if link_desc.connect_retry_interval_ms is not None:
overrides["connect_retry_interval_ms"] = link_desc.connect_retry_interval_ms

if not overrides:
return self._brpc_config
return replace(self._brpc_config, **overrides)

def _get_template_link(
self,
cache_key: tuple,
Expand All @@ -138,6 +172,7 @@ def _get_template_link(
communicator: object | None,
parties: list[int] | None,
spu_endpoints: list[str] | None,
brpc_config: BrpcLinkConfig | None = None,
) -> libspu.link.Context:
"""Get or create a template link for the given configuration.

Expand All @@ -147,7 +182,8 @@ def _get_template_link(

def _create() -> libspu.link.Context:
if spu_endpoints:
return self._create_brpc_link(local_rank, spu_endpoints)
cfg = brpc_config or self._brpc_config
return self._create_brpc_link(local_rank, spu_endpoints, cfg)
elif communicator is not None:
if parties is None:
raise ValueError("parties required when using communicator")
Expand Down Expand Up @@ -203,14 +239,18 @@ def get_or_create(
else:
link_mode = "mem"

brpc_config = (
self._effective_brpc_config(config) if link_mode == "brpc" else None
)
cache_key = (
local_rank,
spu_world_size,
config.protocol,
config.field,
config.fxp_fraction_bits,
link_mode,
tuple(spu_endpoints) if spu_endpoints else None,
self._brpc_config if link_mode == "brpc" else None,
brpc_config,
)

if cache_key in self._runtimes:
Expand All @@ -224,6 +264,7 @@ def get_or_create(
communicator,
parties,
spu_endpoints,
brpc_config,
)
link = template_link.spawn()

Expand Down Expand Up @@ -290,10 +331,13 @@ def _create_channels_link(
return libspu.link.create_with_channels(desc, local_rank, channels)

def _create_brpc_link(
self, local_rank: int, spu_endpoints: list[str]
self,
local_rank: int,
spu_endpoints: list[str],
brpc_config: BrpcLinkConfig | None = None,
) -> libspu.link.Context:
"""Create BRPC link for distributed execution."""
cfg = self._brpc_config
cfg = brpc_config or self._brpc_config
desc = libspu.link.Desc() # type: ignore
desc.recv_timeout_ms = cfg.recv_timeout_ms
desc.http_max_payload_size = cfg.http_max_payload_size
Expand Down
78 changes: 69 additions & 9 deletions mplang/dialects/spu.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,48 @@ def secure_add(x, y):
# ==============================================================================


@dataclass(frozen=True)
class SPULinkDesc:
"""SPU brpc link configuration compatible with SecretFlow link_desc."""

recv_timeout_ms: int | None = None
http_timeout_ms: int | None = None
http_max_payload_size: int | None = None
brpc_channel_protocol: str | None = None
brpc_channel_connection_type: str | None = None
connect_retry_times: int | None = None
connect_retry_interval_ms: int | None = None

@classmethod
def from_dict(cls, d: dict[str, Any] | SPULinkDesc) -> SPULinkDesc:
if isinstance(d, SPULinkDesc):
return d
if not isinstance(d, dict):
raise TypeError(f"SPULinkDesc.from_dict expects dict, got {type(d)!r}")
return cls(
Comment thread
FollyCoolly marked this conversation as resolved.
recv_timeout_ms=d.get("recv_timeout_ms"),
http_timeout_ms=d.get("http_timeout_ms"),
http_max_payload_size=d.get("http_max_payload_size"),
brpc_channel_protocol=d.get("brpc_channel_protocol") or d.get("protocol"),
brpc_channel_connection_type=d.get("brpc_channel_connection_type")
or d.get("connection_type"),
connect_retry_times=d.get("connect_retry_times"),
connect_retry_interval_ms=d.get("connect_retry_interval_ms"),
)

def to_json(self) -> dict[str, Any]:
data = {
"recv_timeout_ms": self.recv_timeout_ms,
"http_timeout_ms": self.http_timeout_ms,
"http_max_payload_size": self.http_max_payload_size,
"brpc_channel_protocol": self.brpc_channel_protocol,
"brpc_channel_connection_type": self.brpc_channel_connection_type,
"connect_retry_times": self.connect_retry_times,
"connect_retry_interval_ms": self.connect_retry_interval_ms,
}
return {key: value for key, value in data.items() if value is not None}


@serde.register_class
@dataclass(frozen=True)
class SPUConfig:
Expand All @@ -103,37 +145,55 @@ class SPUConfig:
protocol: SPU protocol (e.g., "SEMI2K", "ABY3").
field: SPU field type (e.g., "FM64", "FM128").
fxp_fraction_bits: Fixed-point fraction bits.
link_desc: Optional brpc link configuration.
"""

protocol: str = "SEMI2K"
field: str = "FM128"
fxp_fraction_bits: int = 18
link_desc: SPULinkDesc | None = None

@classmethod
def from_dict(cls, d: dict[str, Any]) -> SPUConfig:
if not isinstance(d, dict):
raise TypeError(f"SPUConfig.from_dict expects dict, got {type(d)!r}")
runtime_config = d.get("runtime_config") or {}
if not isinstance(runtime_config, dict):
Comment thread
FollyCoolly marked this conversation as resolved.
raise TypeError(
f"SPUConfig.runtime_config must be a dict, got {runtime_config!r}"
)
link_desc = d.get("link_desc")
protocol = cast(
str, runtime_config.get("protocol", d.get("protocol", "SEMI2K"))
)
field = cast(str, runtime_config.get("field", d.get("field", "FM128")))
fxp_fraction_bits = cast(
int,
runtime_config.get("fxp_fraction_bits", d.get("fxp_fraction_bits", 18)),
)
return cls(
protocol=d.get("protocol", "SEMI2K"),
field=d.get("field", "FM128"),
fxp_fraction_bits=d.get("fxp_fraction_bits", 18),
protocol=protocol,
field=field,
fxp_fraction_bits=fxp_fraction_bits,
link_desc=SPULinkDesc.from_dict(link_desc) if link_desc else None,
)

# --- Serde methods ---
_serde_kind: ClassVar[str] = "spu.SPUConfig"

def to_json(self) -> dict[str, Any]:
return {
data: dict[str, Any] = {
"protocol": self.protocol,
"field": self.field,
"fxp_fraction_bits": self.fxp_fraction_bits,
}
if self.link_desc is not None:
data["link_desc"] = self.link_desc.to_json()
return data

@classmethod
def from_json(cls, data: dict[str, Any]) -> SPUConfig:
return cls(
protocol=data["protocol"],
field=data["field"],
fxp_fraction_bits=data["fxp_fraction_bits"],
)
return cls.from_dict(data)


# ==============================================================================
Expand Down
97 changes: 97 additions & 0 deletions tests/backends/test_spu_state.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
# Copyright 2026 Ant Group Co., Ltd.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Tests for SPU dialect state."""

from __future__ import annotations

from types import SimpleNamespace

from mplang.backends import spu_state
from mplang.backends.spu_state import BrpcLinkConfig, SPUState
from mplang.dialects import spu


def test_effective_brpc_config_merges_spu_link_desc() -> None:
base = BrpcLinkConfig(
protocol="h2",
connection_type="single",
recv_timeout_ms=1000,
http_max_payload_size=1024,
http_timeout_ms=None,
connect_retry_times=3,
connect_retry_interval_ms=4,
)
state = SPUState(brpc_config=base)
config = spu.SPUConfig(
link_desc=spu.SPULinkDesc(
brpc_channel_protocol="http",
brpc_channel_connection_type="pooled",
recv_timeout_ms=7200000,
http_timeout_ms=7200000,
)
)

effective = state._effective_brpc_config(config)

assert effective.protocol == "http"
assert effective.connection_type == "pooled"
assert effective.recv_timeout_ms == 7200000
assert effective.http_timeout_ms == 7200000
assert effective.http_max_payload_size == 1024
assert effective.connect_retry_times == 3
assert effective.connect_retry_interval_ms == 4


def test_create_brpc_link_applies_explicit_config(monkeypatch) -> None:
class FakeDesc:
def __init__(self) -> None:
self.parties: list[tuple[str, str]] = []

def add_party(self, name: str, endpoint: str) -> None:
self.parties.append((name, endpoint))

captured = {}

def fake_create_brpc(desc: FakeDesc, local_rank: int) -> str:
captured["desc"] = desc
captured["local_rank"] = local_rank
return "fake-link"

fake_link = SimpleNamespace(Desc=FakeDesc, create_brpc=fake_create_brpc)
monkeypatch.setattr(spu_state.libspu, "link", fake_link)

cfg = BrpcLinkConfig(
protocol="http",
connection_type="pooled",
recv_timeout_ms=7200000,
http_max_payload_size=67108864,
http_timeout_ms=7200000,
connect_retry_times=9,
connect_retry_interval_ms=10,
)

link = SPUState()._create_brpc_link(1, ["127.0.0.1:9000", "127.0.0.1:9001"], cfg)

assert link == "fake-link"
assert captured["local_rank"] == 1
desc = captured["desc"]
assert desc.recv_timeout_ms == 7200000
assert desc.http_max_payload_size == 67108864
assert desc.http_timeout_ms == 7200000
assert desc.brpc_channel_protocol == "http"
assert desc.brpc_channel_connection_type == "pooled"
assert desc.connect_retry_times == 9
assert desc.connect_retry_interval_ms == 10
assert desc.parties == [("P0", "127.0.0.1:9000"), ("P1", "127.0.0.1:9001")]
Loading
Loading