Skip to content
Open
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
6 changes: 6 additions & 0 deletions modelexpress_client/python/modelexpress/envs.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,8 @@
MX_RESHARD_HANDSHAKE_BACKOFF_S: float
MX_REFIT_STAGE_RECORD: bool
MX_RESHARD_MAX_GBPS: float
MX_REFIT_S3_UPLOAD_WORKERS: int
MX_REFIT_S3_MAX_POOL_CONNECTIONS: int
# Kubernetes service backend
MX_K8S_SERVICE_PATTERN: str
MX_K8S_SOURCE_RETRIES: str
Expand Down Expand Up @@ -268,6 +270,10 @@ def _env_positive_float(name: str, default: float) -> float:
# disables the check, and is the default because only the operator knows the
# real per-rank limit for their fabric.
"MX_RESHARD_MAX_GBPS": lambda: _env_float("MX_RESHARD_MAX_GBPS", 0.0),
"MX_REFIT_S3_UPLOAD_WORKERS": lambda: _env_int("MX_REFIT_S3_UPLOAD_WORKERS", 4),
"MX_REFIT_S3_MAX_POOL_CONNECTIONS": lambda: _env_int(
"MX_REFIT_S3_MAX_POOL_CONNECTIONS", 10
),
# ── Kubernetes service backend ─────────────────────────────────────────
"MX_K8S_SERVICE_PATTERN": lambda: os.environ.get("MX_K8S_SERVICE_PATTERN", "mx-sources"),
"MX_K8S_SOURCE_RETRIES": lambda: os.environ.get("MX_K8S_SOURCE_RETRIES", ""),
Expand Down
22 changes: 22 additions & 0 deletions modelexpress_client/python/modelexpress/refit/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,16 @@

"""Engine-agnostic primitives for live model refit."""

from .api import (
PublisherConfig,
ReceiverRevisionState,
ReceiverStatus,
S3Config,
WeightUpdateResult,
)
from .catalog import GrpcRevisionCatalog, RevisionCatalog
from .manifest import RevisionManifest, RevisionRecord, RevisionState, S3Object
from .publisher import Publisher
from .timing import (
MX_REFIT_TIMING_PREFIX,
REFIT_TIMING_STAGES,
Expand All @@ -14,9 +24,21 @@
)

__all__ = [
"GrpcRevisionCatalog",
"MX_REFIT_TIMING_PREFIX",
"Publisher",
"PublisherConfig",
"REFIT_TIMING_STAGES",
"ReceiverRevisionState",
"ReceiverStatus",
"RefitTimingRecorder",
"RevisionCatalog",
"RevisionManifest",
"RevisionRecord",
"RevisionState",
"S3Config",
"S3Object",
"WeightUpdateResult",
"add_refit_bytes",
"current_refit_timing",
"refit_span",
Expand Down
62 changes: 62 additions & 0 deletions modelexpress_client/python/modelexpress/refit/api.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Values consumed by the Miles publisher and SGLang receiver."""

from __future__ import annotations

from dataclasses import dataclass
from enum import Enum

ModelId = str
VersionId = str


class ReceiverRevisionState(Enum):
"""SGLang-local receiver outcomes, never persisted by the MX server."""

BYTES_RECEIVED = "bytes_received"
VERIFIED = "verified"
FAILED = "failed"
POISONED = "poisoned"


@dataclass(frozen=True)
class S3Config:
"""Direct S3 destination; credentials are resolved privately by boto3."""

bucket: str
prefix: str = ""
endpoint_url: str | None = None
region_name: str | None = None


@dataclass(frozen=True)
class PublisherConfig:
model_id: ModelId
catalog_endpoint: str
s3: S3Config


@dataclass(frozen=True)
class WeightUpdateResult:
success: bool
receiver_id: str
installed_version: VersionId | None
state: ReceiverRevisionState
target_digest: str | None = None
detail: str = ""


@dataclass(frozen=True)
class ReceiverStatus:
receiver_id: str
model_id: ModelId
installed_version: VersionId | None = None
target_digest: str | None = None
state: ReceiverRevisionState | None = None
detail: str = ""

@property
def recovery_required(self) -> bool:
return self.state is ReceiverRevisionState.POISONED
108 changes: 108 additions & 0 deletions modelexpress_client/python/modelexpress/refit/delta.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Canonical XOR-delta bucket framing and zstd compression."""

from __future__ import annotations

import hashlib
import json
import struct

import numpy as np
import zstandard

from .source.canonical import canonical_json

_BUCKET_MAGIC = b"MXCDV0\0"
_SCHEMA = "mx.canonical.delta.v0"


def compute_delta(
current: np.ndarray, base: np.ndarray
) -> tuple[np.ndarray | None, str | None, int]:
if len(current) != len(base):
raise RuntimeError("tensor changed byte size")
raw_delta = np.bitwise_xor(current, base)
changed_bytes = int(np.count_nonzero(raw_delta))
if not changed_bytes:
return None, None, 0
target_digest = f"sha256:{hashlib.sha256(memoryview(current)).hexdigest()}"
return raw_delta, target_digest, changed_bytes


def encode_bucket(
model_id: str,
base_version: str,
target_version: str,
base_digest: str,
format_digest: str,
ordinal: int,
tensors: list[tuple[str, np.ndarray]],
metadata: dict[str, dict],
) -> tuple[bytes, int]:
entries = []
offset = 0
compressor = zstandard.ZstdCompressor(level=3).compressobj()
chunks = []
for name, delta in tensors:
entries.append({**metadata[name], "offset": offset})
offset += delta.nbytes
chunks.append(compressor.compress(memoryview(delta)))
chunks.append(compressor.flush())

header = canonical_json(
{
"base_digest": base_digest,
"base_version": base_version,
"compression": "zstd",
"decoded_size": offset,
"delta": "xor",
"entries": entries,
"format_digest": format_digest,
"model_id": model_id,
"ordinal": ordinal,
"schema": f"{_SCHEMA}.bucket",
"target_version": target_version,
}
)
compressed = b"".join(chunks)
return _BUCKET_MAGIC + struct.pack(">I", len(header)) + header + compressed, offset


def bucket_parts(data: bytes) -> tuple[dict, memoryview]:
if not data.startswith(_BUCKET_MAGIC):
raise ValueError("invalid canonical bucket")
header_size = struct.unpack(
">I", data[len(_BUCKET_MAGIC) : len(_BUCKET_MAGIC) + 4]
)[0]
header_start = len(_BUCKET_MAGIC) + 4
header = json.loads(data[header_start : header_start + header_size])
return header, memoryview(data)[header_start + header_size :]


def parse_bucket(data: bytes) -> tuple[dict, bytes]:
header, compressed = bucket_parts(data)
decoded = zstandard.ZstdDecompressor().decompress(
compressed, max_output_size=header["decoded_size"]
)
return header, decoded


def decode_bucket(
data: bytes, snapshot: dict[str, np.ndarray], metadata: dict[str, dict]
) -> dict:
header, decoded = parse_bucket(data)
for entry in header["entries"]:
name = entry["name"]
start = entry["offset"]
delta = np.frombuffer(
decoded[start : start + entry["byte_size"]], dtype=np.uint8
)
target = np.bitwise_xor(snapshot[name], delta)
digest = f"sha256:{hashlib.sha256(target.tobytes()).hexdigest()}"
if digest != entry["target_digest"]:
raise ValueError(f"canonical target checksum differs for {name}")
snapshot[name] = target
metadata[name]["target_digest"] = digest
return header
Loading
Loading