Skip to main content

Zero-copy model sharing for PyTorch inference in Ray

Project description

ray-zerocopy

PyPI version Documentation Status License

Zero-copy model sharing for PyTorch inference in Ray

This library enables efficient model sharing across Ray workers using zero-copy mechanisms, eliminating the need to duplicate large model weights in memory when performing inference.

Features

  • 🚀 Zero-copy sharing - Share model weights across Ray workers without duplication
  • 🎯 Flexible inference - Use with Ray Tasks, Ray Actors, or Ray Data Actor UDFs
  • 💾 Memory efficient - 4 actors with 5GB model = ~5GB total (not 20GB)
  • High throughput - Direct inference without model loading overhead
  • 🔧 Pipeline support - Share entire pipelines (classes with nn.Module attributes)

Quick Start

For Ray Data Actor UDFs (Recommended for Batch Inference)

from ray.data import ActorPoolStrategy
from ray_zerocopy import ModelWrapper

# 1. Create your pipeline (a class with nn.Module attributes)
class MyPipeline:
    def __init__(self):
        self.encoder = EncoderModel()
        self.decoder = DecoderModel()

    def __call__(self, data):
        encoded = self.encoder(data)
        return self.decoder(encoded)

pipeline = MyPipeline()

# 2. Wrap with ModelWrapper for zero-copy sharing
model_wrapper = ModelWrapper.from_model(pipeline, mode="actor")

# 3. Define actor UDF that loads the pipeline
class InferenceActor:
    def __init__(self, model_wrapper):
        self.pipeline = model_wrapper.load()

    def __call__(self, batch):
        with torch.no_grad():
            return self.pipeline(batch["data"])

# 4. Use with Ray Data
results = ds.map_batches(
    InferenceActor,
    fn_constructor_kwargs={"model_wrapper": model_wrapper},
    compute=ActorPoolStrategy(size=4),  # 4 actors share the model
)

For Ray Actors (General Purpose)

import ray
from ray_zerocopy import ModelWrapper

# Wrap pipeline for actors
pipeline = MyPipeline()
model_wrapper = ModelWrapper.from_model(pipeline, mode="actor")

# Define inference actor
@ray.remote
class InferenceActor:
    def __init__(self, model_wrapper):
        self.pipeline = model_wrapper.load()

    def predict(self, data):
        with torch.no_grad():
            return self.pipeline(data)

# Create actors that share the model
actors = [InferenceActor.remote(model_wrapper) for _ in range(4)]
results = ray.get([actor.predict.remote(data) for actor in actors])

For Ray Tasks (Ad-hoc Inference)

from ray_zerocopy import ModelWrapper

# A Pipeline is a class with nn.Module attributes
class MyPipeline:
    def __init__(self):
        self.encoder = EncoderModel()
        self.decoder = DecoderModel()

    def __call__(self, data):
        encoded = self.encoder(data)
        return self.decoder(encoded)

pipeline = MyPipeline()
wrapped = ModelWrapper.for_tasks(pipeline)

# Each call spawns a Ray task with zero-copy model loading
result = wrapped(data)

Installation

pip install ray-zerocopy

Or install from source:

git clone https://github.com/wingkitlee0/ray-zerocopy.git
cd ray-zerocopy
pip install -e .

When to Use What

Scenario Use This
Ray Data map_batches batch inference ModelWrapper.from_model(..., mode="actor") with Ray Data Actor UDF
High-throughput batch inference ModelWrapper.from_model(..., mode="actor") with Ray Data Actor UDF
Long-running inference service ModelWrapper.from_model(..., mode="actor") with Ray Actor
Ad-hoc task-based inference ModelWrapper.for_tasks() with Ray Task
Sporadic inference calls ModelWrapper.for_tasks() with Ray Task

Memory Savings Example

Without zero-copy:

Actor 1: 5GB model
Actor 2: 5GB model
Actor 3: 5GB model
Actor 4: 5GB model
Total: 20GB

With zero-copy:

Ray Object Store: 5GB (shared)
Actor 1-4: reference object store
Total: ~5GB

Pipelines

A Pipeline is a class with nn.Module attributes. The library automatically identifies and shares all models in a pipeline:

class MyPipeline:
    def __init__(self):
        self.feature_extractor = FeatureExtractorModel()
        self.classifier = ClassifierModel()
        self.config = {"threshold": 0.5}  # Non-model attributes are preserved

    def __call__(self, data):
        features = self.feature_extractor(data)
        return self.classifier(features)

# For Ray Actors and Ray Data
model_wrapper = ModelWrapper.from_model(pipeline, mode="actor")
# ... in actor:
# self.pipeline = model_wrapper.load()

# For Ray Tasks
wrapped = ModelWrapper.for_tasks(pipeline)

The library automatically identifies nn.Module attributes and applies zero-copy sharing to them, while preserving other attributes like config dictionaries.

API Overview

Wrapper Classes

from ray_zerocopy import ModelWrapper

# ModelWrapper - For Ray Tasks
wrapped = ModelWrapper.for_tasks(pipeline)
result = wrapped(data)  # Runs in Ray task with zero-copy

# ModelWrapper - For Ray Actors and Ray Data
model_wrapper = ModelWrapper.from_model(pipeline, mode="actor")
# ... in actor __init__:
pipeline = model_wrapper.load()  # Load with zero-copy in actor

TorchScript Support

from ray_zerocopy import JITTaskWrapper, JITActorWrapper

# JITTaskWrapper - For TorchScript models with Ray Tasks
jit_pipeline = torch.jit.trace(pipeline, example_input)
wrapped = JITTaskWrapper(jit_pipeline)

# JITActorWrapper - For TorchScript models with Ray Actors
actor_wrapper = JITActorWrapper(jit_pipeline)

Requirements

  • Python 3.8+
  • PyTorch
  • Ray

Acknowledgments

This project includes code derived from IBM's Zero-Copy Model Loading project, licensed under Apache 2.0.

License

Apache License 2.0 (see LICENSE)

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

ray_zerocopy-0.1.0.tar.gz (74.0 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

ray_zerocopy-0.1.0-py3-none-any.whl (36.2 kB view details)

Uploaded Python 3

File details

Details for the file ray_zerocopy-0.1.0.tar.gz.

File metadata

  • Download URL: ray_zerocopy-0.1.0.tar.gz
  • Upload date:
  • Size: 74.0 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.7

File hashes

Hashes for ray_zerocopy-0.1.0.tar.gz
Algorithm Hash digest
SHA256 686166b9c49966d30c5d79bac21766945af9a9ec2085af1363f4d6fd8d8916d8
MD5 768d1029303e2c0171e890cac5d27ce4
BLAKE2b-256 a0a622f9ec0484628e739c2dd793d1ba8188aaccf8e48e4b53e63e7873a64917

See more details on using hashes here.

Provenance

The following attestation bundles were made for ray_zerocopy-0.1.0.tar.gz:

Publisher: publish.yml on wingkitlee0/ray-zerocopy

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file ray_zerocopy-0.1.0-py3-none-any.whl.

File metadata

  • Download URL: ray_zerocopy-0.1.0-py3-none-any.whl
  • Upload date:
  • Size: 36.2 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.7

File hashes

Hashes for ray_zerocopy-0.1.0-py3-none-any.whl
Algorithm Hash digest
SHA256 e0c4813b7a1f4b79215e3809b00541a699ad9b15dd7ba2d1f6d6e89ef03af274
MD5 7ce662d56ae015f73dfd66d4718a4148
BLAKE2b-256 16b27b75d8b98ea2fe0f7c8bfb5c70a44a14ea6d013223ce74907fd35caafc26

See more details on using hashes here.

Provenance

The following attestation bundles were made for ray_zerocopy-0.1.0-py3-none-any.whl:

Publisher: publish.yml on wingkitlee0/ray-zerocopy

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page