Files
2026-07-13 13:17:40 +08:00

259 lines
8.1 KiB
Python

import argparse
import io
import uuid
from typing import Any, Dict
import numpy as np
import pandas as pd
import torch
from benchmark import (
Benchmark,
RuntimeEnvSetupTracker,
benchmark_py_modules,
collect_dataset_stats,
)
from PIL import Image
from torchvision.models import vit_b_16, ViT_B_16_Weights
import albumentations as A
import ray
import copy
import itertools
from typing import List
import string
import random
import time
from ray.data.expressions import download
from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy
from ray._private.test_utils import EC2InstanceTerminatorWithGracePeriod
WRITE_PATH = f"s3://ray-data-write-benchmark/{uuid.uuid4().hex}"
BUCKET = "ray-benchmark-data-internal-us-west-2"
# Assumptions: homogenously shaped images, homogenous images
# Each image is 2048 * 2048 * 3 = 12.58 MB -> 11 images / block. 8 blocks per task, so ~88 images per task.
IMAGES_PER_BLOCK = 11
BLOCKS_PER_TASK = 8
NUM_UNITS = 1380
NUM_CONTAINERS = 50
OVERRIDE_NUM_BLOCKS = int(NUM_CONTAINERS * NUM_UNITS / IMAGES_PER_BLOCK)
PATCH_SIZE = 256
# Largest batch that can fit on a T4.
BATCH_SIZE = 1200
# On a T4 GPU, it takes ~11.3s to perform inference on 1200 images. So, the time per
# image is 11.3s / 1200 ~= 0.0094s.
INFERENCE_LATENCY_PER_IMAGE_S = 0.0094
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument(
"--inference-concurrency",
nargs=2,
type=int,
required=True,
help="The minimum and maximum concurrency for the inference operator.",
)
parser.add_argument(
"--sf",
dest="scale_factor",
type=int,
default=1,
help=(
"The number of copies of the dataset to read. Use this to simulate a larger "
"dataset."
),
)
parser.add_argument(
"--chaos",
action="store_true",
help=(
"Whether to enable chaos. If set, this script terminates one worker node "
"every minute with a grace period."
),
)
return parser.parse_args()
def create_metadata(scale_factor: int):
# TODO(mowen): Handle repeats of the dataset if scale_factor > 1
# simulate various text metadata fields alongside image metadata
return pd.DataFrame(
[
{
"metadata_0": "".join(random.choices(string.ascii_letters, k=16)),
"metadata_1": "".join(random.choices(string.ascii_letters, k=16)),
"metadata_2": "".join(random.choices(string.ascii_letters, k=16)),
"metadata_3": "".join(random.choices(string.ascii_letters, k=16)),
"metadata_4": "".join(random.choices(string.ascii_letters, k=16)),
"metadata_5": "".join(random.choices(string.ascii_letters, k=16)),
"metadata_6": "".join(random.choices(string.ascii_letters, k=16)),
"container_order_read_id": f"{i:04d}_{j:04d}",
"container_id": i,
"channel0_uris": f"s3://{BUCKET}/15TiB-high-resolution-images/group={i:04d}/{j:04d}_{0}.png",
"channel1_uris": f"s3://{BUCKET}/15TiB-high-resolution-images/group={i:04d}/{j:04d}_{1}.png",
"channel2_uris": f"s3://{BUCKET}/15TiB-high-resolution-images/group={i:04d}/{j:04d}_{2}.png",
"applied_scale": 1,
}
for j in range(NUM_UNITS)
for i in range(NUM_CONTAINERS)
]
)
def combine_channels(row: Dict[str, Any]) -> Dict[str, np.ndarray]:
channels = []
for i in range(3):
data = io.BytesIO(row.pop(f"channel{i}"))
image = Image.open(data)
channels.append(np.array(image))
row["image"] = np.dstack(channels)
return row
def process_image(row: Dict[str, Any]) -> Dict[str, np.ndarray]:
transform = A.Compose(
[
A.ToFloat(),
A.LongestMaxSize(
max_size=int(row["image"].shape[0] * float(1.0 / row["applied_scale"]))
),
A.FromFloat(dtype="uint8"),
]
)
row["image"] = transform(image=row["image"])["image"]
return row
def patch_image(row: Dict[str, Any]) -> List[Dict[str, Any]]:
image = row.pop("image")
patches = []
width, height, _ = image.shape
for x, y in itertools.product(
range(PATCH_SIZE, width - PATCH_SIZE, PATCH_SIZE),
range(PATCH_SIZE, height - PATCH_SIZE, PATCH_SIZE),
):
patch = image[y : y + PATCH_SIZE, x : x + PATCH_SIZE, :]
patch_row = copy.deepcopy(row)
patch_row["patch_x"] = x
patch_row["patch_y"] = y
patch_row["patch_width"] = PATCH_SIZE
patch_row["patch_height"] = PATCH_SIZE
patch_row["patch"] = patch
patches.append(patch_row)
return patches
class ProcessPatches:
def __init__(self, transform):
self._transform = transform
def __call__(self, batch: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
batch["patch"] = self._transform(
torch.as_tensor(batch["patch"]).permute(0, 3, 1, 2)
)
return batch
class EmbedPatches:
def __init__(self, model, device):
self._model = ray.get(model)
self._model.eval()
self._model.to(device)
self._device = device
def __call__(self, batch: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
inputs = torch.as_tensor(batch.pop("patch"), device=self._device)
with torch.inference_mode():
output = self._model(inputs)
batch["embedding"] = output.cpu().numpy()
return batch
class FakeEmbedPatches:
def __init__(self, model, device):
self._model = ray.get(model)
self._model.eval()
def __call__(self, batch: Dict[str, np.ndarray]) -> Dict[str, np.ndarray]:
inputs = torch.as_tensor(batch.pop("patch"))
with torch.inference_mode():
# Simulate inference latency with a sleep
time.sleep(INFERENCE_LATENCY_PER_IMAGE_S * len(inputs))
# Generate fake embeddings
output = torch.rand((len(inputs), 1000), dtype=torch.float)
batch["embedding"] = output.cpu().numpy()
return batch
def main(args: argparse.Namespace):
benchmark = Benchmark()
if args.chaos:
start_chaos()
print("Creating metadata")
metadata = create_metadata(scale_factor=args.scale_factor)
def benchmark_fn():
weights = ViT_B_16_Weights.DEFAULT
model = vit_b_16(weights=weights)
transform = weights.transforms()
model_ref = ray.put(model)
ds = (
ray.data.from_pandas(metadata)
.with_column("channel0", download("channel0_uris"))
.with_column("channel1", download("channel1_uris"))
.with_column("channel2", download("channel2_uris"))
.map(combine_channels)
.filter(lambda row: row["image"].size != 0)
.map(process_image)
.flat_map(patch_image)
.map_batches(ProcessPatches(transform), batch_size="auto")
.map_batches(
EmbedPatches,
num_gpus=1,
batch_size=BATCH_SIZE,
concurrency=tuple(args.inference_concurrency),
fn_constructor_kwargs={"model": model_ref, "device": "cuda"},
)
)
ds.write_parquet(WRITE_PATH)
metrics = collect_dataset_stats(ds)
metrics["runtime_env_setup"] = RuntimeEnvSetupTracker.collect()
return metrics
benchmark.run_fn("main", benchmark_fn)
benchmark.write_result()
def start_chaos():
assert ray.is_initialized()
head_node_id = ray.get_runtime_context().get_node_id()
scheduling_strategy = NodeAffinitySchedulingStrategy(
node_id=head_node_id, soft=False
)
resource_killer = EC2InstanceTerminatorWithGracePeriod.options(
scheduling_strategy=scheduling_strategy
).remote(head_node_id, max_to_kill=None)
ray.get(resource_killer.ready.remote())
resource_killer.run.remote()
if __name__ == "__main__":
args = parse_args()
ray.init(runtime_env={"py_modules": benchmark_py_modules()})
main(args)