From d531641b07ee9ada41a5ebc9f171b4a62a61d95f Mon Sep 17 00:00:00 2001 From: Daniel Wirjo Date: Sat, 29 Aug 2026 20:46:18 -0700 Subject: [PATCH 1/3] Add vLLM-Omni image-to-video sample --- 03-features/vllm-omni-image-video/.gitignore | 4 + 03-features/vllm-omni-image-video/README.md | 151 ++++++ 03-features/vllm-omni-image-video/app.py | 138 +++++ 03-features/vllm-omni-image-video/cleanup.py | 41 ++ 03-features/vllm-omni-image-video/deploy.py | 123 +++++ 03-features/vllm-omni-image-video/generate.py | 121 +++++ .../requirements-dev.txt | 4 + .../vllm-omni-image-video/requirements.txt | 3 + .../tests/test_vllm_omni_media.py | 256 ++++++++++ .../vllm-omni-image-video/vllm_omni_media.py | 476 ++++++++++++++++++ 10 files changed, 1317 insertions(+) create mode 100644 03-features/vllm-omni-image-video/.gitignore create mode 100644 03-features/vllm-omni-image-video/README.md create mode 100644 03-features/vllm-omni-image-video/app.py create mode 100644 03-features/vllm-omni-image-video/cleanup.py create mode 100644 03-features/vllm-omni-image-video/deploy.py create mode 100644 03-features/vllm-omni-image-video/generate.py create mode 100644 03-features/vllm-omni-image-video/requirements-dev.txt create mode 100644 03-features/vllm-omni-image-video/requirements.txt create mode 100644 03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py create mode 100644 03-features/vllm-omni-image-video/vllm_omni_media.py diff --git a/03-features/vllm-omni-image-video/.gitignore b/03-features/vllm-omni-image-video/.gitignore new file mode 100644 index 0000000..396dc9b --- /dev/null +++ b/03-features/vllm-omni-image-video/.gitignore @@ -0,0 +1,4 @@ +.vllm_omni_media_state.json +.pytest_cache/ +.ruff_cache/ +outputs/ diff --git a/03-features/vllm-omni-image-video/README.md b/03-features/vllm-omni-image-video/README.md new file mode 100644 index 0000000..39d5f80 --- /dev/null +++ b/03-features/vllm-omni-image-video/README.md @@ -0,0 +1,151 @@ +# Generate images and video with vLLM-Omni on SageMaker AI + +This sample deploys two diffusion models with the +[AWS Deep Learning Container for vLLM-Omni](https://aws.github.io/deep-learning-containers/vllm-omni/): + +- [FLUX.2-klein-4B](https://huggingface.co/black-forest-labs/FLUX.2-klein-4B) + generates an image through a SageMaker real-time endpoint. +- [Wan2.1-VACE-1.3B](https://huggingface.co/Wan-AI/Wan2.1-VACE-1.3B-diffusers) + animates that image through SageMaker Asynchronous Inference. + +The command-line workflow and Streamlit application use the same deployment +state and request helpers. + +## Architecture + +```mermaid +flowchart LR + User[CLI or Streamlit app] --> ImageEndpoint[FLUX.2-klein
real-time endpoint] + ImageEndpoint --> PNG[Generated PNG] + PNG --> Request[Multipart video request] + Request --> S3Input[Amazon S3 input] + S3Input --> VideoEndpoint[Wan VACE
asynchronous endpoint] + VideoEndpoint --> S3Output[Amazon S3 MP4 output] + S3Output --> User +``` + +The image endpoint returns an OpenAI-compatible JSON response containing a +base64-encoded PNG. The sample places that PNG in the video request as a base64 +data URL. It uploads the complete multipart request to Amazon S3, invokes +`/v1/videos/sync` through the asynchronous endpoint, and downloads the MP4 from +the returned output location. + +## Prerequisites + +- An AWS account with permissions to create SageMaker models, endpoint + configurations, endpoints, and Amazon S3 objects. +- A SageMaker execution role that can read from the AWS Deep Learning Container + registry and access the sample Amazon S3 bucket. +- Endpoint quota for one `ml.g6.xlarge` instance and one `ml.g6e.xlarge` + instance in the selected AWS Region. +- Python 3.11 or later. + +This sample pins `omni-sagemaker-cuda-v1.6`, which contains vLLM-Omni 0.26.0. +Review the +[vLLM-Omni DLC changelog](https://aws.github.io/deep-learning-containers/vllm-omni/changelog/) +before changing the tag. + +## Install the dependencies + +Create a virtual environment and install the sample dependencies: + +```bash +python -m venv .venv +source .venv/bin/activate +python -m pip install -r requirements.txt +``` + +## Deploy the endpoints + +Set the SageMaker execution role, then deploy both models in `us-east-1`: + +```bash +export SAGEMAKER_ROLE_ARN= + +python deploy.py \ + --role-arn "$SAGEMAKER_ROLE_ARN" \ + --region us-east-1 +``` + +The deployment script writes resource names to +`.vllm_omni_media_state.json`. Keep this file until you finish generation and +cleanup. + +FLUX.2-klein uses a real-time endpoint because its response completes within a +single invocation. Wan VACE uses SageMaker Asynchronous Inference because video +generation can exceed the response window for real-time inference. The +asynchronous endpoint is configured for one concurrent request per instance. + +## Generate an image and video + +Run the complete workflow: + +```bash +python generate.py \ + --image-prompt "Cinematic photograph of a coastal observatory at sunrise" \ + --video-prompt "Slow camera push-in as clouds move and grass bends in the wind" +``` + +The script writes the PNG and MP4 to `outputs/`. It also prints the Amazon S3 +output location returned by SageMaker Asynchronous Inference. + +The default video settings use 17 frames and four diffusion steps to keep a +validation run short. Increase `--frames` and `--video-steps` after you confirm +the deployment, then assess output quality and instance memory for your +workload. + +## Run the Streamlit application + +Start the local application after both endpoints reach `InService`: + +```bash +streamlit run app.py +``` + +Generate the source image first, review it, then enter a motion prompt and +generate the video. The application can display and download both artifacts. + +## Request routing + +SageMaker sends inference traffic to `/invocations`. The vLLM-Omni SageMaker +image reads `CustomAttributes` and forwards the request to the selected +OpenAI-compatible route: + +| Workload | Route | SageMaker inference option | +| --- | --- | --- | +| Image generation | `/v1/images/generations` | Real-time inference | +| Video generation | `/v1/videos/sync` | Asynchronous inference | + +The Videos API accepts multipart form data. This sample pre-builds the +multipart body before uploading it to Amazon S3, which keeps the request format +explicit and works across vLLM-Omni DLC releases that accept multipart +requests. + +## Clean up + +Delete the endpoints, endpoint configurations, and models: + +```bash +python cleanup.py +``` + +The cleanup script retains generated objects under the Amazon S3 output prefix +so you can review them. Delete that prefix when you no longer need the +artifacts. + +## Production considerations + +Treat this sample as a deployment baseline, not a production architecture. +Before serving application traffic: + +- Apply least-privilege IAM policies to the deployment identity and SageMaker + execution role. +- Place endpoint traffic in your network design and configure encryption keys, + logging, monitoring, and retention policies for your requirements. +- Load test each model separately, then set instance types, concurrency, and + autoscaling from measured latency, memory, and throughput. +- Add input policy controls and output review appropriate for generated media. + +Review +[SageMaker Asynchronous Inference](https://docs.aws.amazon.com/sagemaker/latest/dg/async-inference.html) +for payload, processing, scaling, and notification options. diff --git a/03-features/vllm-omni-image-video/app.py b/03-features/vllm-omni-image-video/app.py new file mode 100644 index 0000000..7e94fa4 --- /dev/null +++ b/03-features/vllm-omni-image-video/app.py @@ -0,0 +1,138 @@ +"""Streamlit interface for the vLLM-Omni image-to-video workflow.""" + +from __future__ import annotations + +import time +from pathlib import Path + +import boto3 +import streamlit as st + +from vllm_omni_media import ( + DEFAULT_STATE_PATH, + invoke_image, + load_state, + submit_video, + validate_mp4, + wait_for_s3_object, +) + +st.set_page_config(page_title="vLLM-Omni image to video", layout="wide") +st.title("Generate images and video with vLLM-Omni") + +state_path = Path( + st.sidebar.text_input("Deployment state", str(DEFAULT_STATE_PATH)) +).expanduser() + +try: + state = load_state(state_path) +except (FileNotFoundError, TypeError, ValueError) as error: + st.info(str(error)) + st.stop() + +session = boto3.Session(region_name=state.region) +runtime = session.client("sagemaker-runtime") +s3 = session.client("s3") + +image_prompt = st.text_area( + "Image prompt", + ( + "Cinematic photograph of a small observatory on a windswept coastal " + "cliff at sunrise, detailed clouds, natural light" + ), +) +image_col, seed_col = st.columns([3, 1]) +with image_col: + image_size = st.selectbox("Image size", ["1024x1024", "768x768"]) +with seed_col: + seed = st.number_input("Seed", min_value=0, value=42, step=1) + +if st.button("Generate image", type="primary", use_container_width=True): + with st.spinner("Generating image"): + started = time.perf_counter() + st.session_state.image_bytes = invoke_image( + runtime, + state, + image_prompt, + size=image_size, + seed=int(seed), + ) + st.session_state.image_seconds = time.perf_counter() - started + st.session_state.pop("video_bytes", None) + +if "image_bytes" in st.session_state: + st.image( + st.session_state.image_bytes, + caption=f"FLUX.2-klein output in {st.session_state.image_seconds:.1f}s", + use_container_width=True, + ) + st.download_button( + "Download image", + st.session_state.image_bytes, + "flux2-klein.png", + "image/png", + use_container_width=True, + ) + +st.divider() +video_prompt = st.text_area( + "Motion prompt", + ( + "Slow camera push-in as clouds move across the sky and grass bends " + "in the wind, stable composition" + ), +) +frames_col, fps_col, steps_col = st.columns(3) +with frames_col: + frames = st.select_slider("Frames", options=[9, 17, 33], value=17) +with fps_col: + fps = st.select_slider("Frames per second", options=[4, 8, 16], value=8) +with steps_col: + steps = st.select_slider("Inference steps", options=[4, 8, 16, 30], value=4) + +if st.button( + "Generate video", + disabled="image_bytes" not in st.session_state, + use_container_width=True, +): + with st.status("Generating video", expanded=True) as status: + status.write("Uploading the multipart request to Amazon S3") + output_uri, failure_uri, request_key = submit_video( + runtime, + s3, + state, + video_prompt, + st.session_state.image_bytes, + num_frames=int(frames), + fps=int(fps), + steps=int(steps), + seed=int(seed), + ) + status.write("Waiting for the SageMaker asynchronous endpoint") + started = time.perf_counter() + try: + st.session_state.video_bytes = wait_for_s3_object( + s3, + output_uri, + failure_uri=failure_uri, + ) + finally: + s3.delete_object(Bucket=state.bucket, Key=request_key) + validate_mp4(st.session_state.video_bytes) + st.session_state.video_seconds = time.perf_counter() - started + st.session_state.output_uri = output_uri + status.update(label="Video ready", state="complete", expanded=False) + +if "video_bytes" in st.session_state: + st.video(st.session_state.video_bytes) + st.caption( + f"Wan VACE output in {st.session_state.video_seconds:.1f}s. " + f"Stored at {st.session_state.output_uri}" + ) + st.download_button( + "Download video", + st.session_state.video_bytes, + "wan-vace.mp4", + "video/mp4", + use_container_width=True, + ) diff --git a/03-features/vllm-omni-image-video/cleanup.py b/03-features/vllm-omni-image-video/cleanup.py new file mode 100644 index 0000000..9b4cc59 --- /dev/null +++ b/03-features/vllm-omni-image-video/cleanup.py @@ -0,0 +1,41 @@ +#!/usr/bin/env python3 +"""Delete the SageMaker resources created by deploy.py.""" + +from __future__ import annotations + +import argparse +from pathlib import Path + +import boto3 + +from vllm_omni_media import ( + DEFAULT_STATE_PATH, + delete_deployment, + load_state, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--state-file", + type=Path, + default=DEFAULT_STATE_PATH, + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + state = load_state(args.state_file) + sagemaker = boto3.client("sagemaker", region_name=state.region) + delete_deployment(sagemaker, state) + args.state_file.unlink(missing_ok=True) + print("Deleted the image and video endpoint resources.") + print( + f"Generated artifacts remain under s3://{state.bucket}/{state.prefix}/outputs/." + ) + + +if __name__ == "__main__": + main() diff --git a/03-features/vllm-omni-image-video/deploy.py b/03-features/vllm-omni-image-video/deploy.py new file mode 100644 index 0000000..88f43a9 --- /dev/null +++ b/03-features/vllm-omni-image-video/deploy.py @@ -0,0 +1,123 @@ +#!/usr/bin/env python3 +"""Deploy FLUX.2-klein and Wan VACE with the vLLM-Omni DLC.""" + +from __future__ import annotations + +import argparse +import re +from datetime import UTC, datetime +from pathlib import Path + +import boto3 + +from vllm_omni_media import ( + DEFAULT_STATE_PATH, + IMAGE_MODEL_ID, + VIDEO_MODEL_ID, + DeploymentState, + container_image_uri, + create_endpoint, + ensure_bucket, + save_state, + wait_for_endpoint, +) + + +def resource_name(prefix: str, kind: str, timestamp: str) -> str: + """Build a SageMaker-compatible resource name.""" + + value = re.sub(r"[^A-Za-z0-9-]", "-", f"{prefix}-{kind}-{timestamp}") + return value[:63].rstrip("-") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Deploy vLLM-Omni image and video endpoints on SageMaker AI." + ) + parser.add_argument("--role-arn", required=True) + parser.add_argument("--region", default="us-east-1") + parser.add_argument( + "--bucket", + help="S3 bucket for async requests and outputs. Defaults to the SageMaker convention.", + ) + parser.add_argument("--prefix", default="vllm-omni-image-video") + parser.add_argument("--name-prefix", default="vllm-omni-media") + parser.add_argument("--image-instance-type", default="ml.g6.xlarge") + parser.add_argument("--video-instance-type", default="ml.g6e.xlarge") + parser.add_argument( + "--state-file", + type=Path, + default=DEFAULT_STATE_PATH, + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + session = boto3.Session(region_name=args.region) + sts = session.client("sts") + account_id = sts.get_caller_identity()["Account"] + bucket = args.bucket or f"sagemaker-{args.region}-{account_id}" + sagemaker = session.client("sagemaker") + s3 = session.client("s3") + ensure_bucket(s3, bucket, args.region) + + timestamp = datetime.now(UTC).strftime("%Y%m%d%H%M%S") + state = DeploymentState( + region=args.region, + bucket=bucket, + prefix=args.prefix.strip("/"), + image_model_name=resource_name(args.name_prefix, "image-model", timestamp), + image_endpoint_config_name=resource_name( + args.name_prefix, "image-config", timestamp + ), + image_endpoint_name=resource_name( + args.name_prefix, "image-endpoint", timestamp + ), + video_model_name=resource_name(args.name_prefix, "video-model", timestamp), + video_endpoint_config_name=resource_name( + args.name_prefix, "video-config", timestamp + ), + video_endpoint_name=resource_name( + args.name_prefix, "video-endpoint", timestamp + ), + ) + save_state(state, args.state_file) + + image_uri = container_image_uri(args.region) + print(f"Creating image endpoint {state.image_endpoint_name}") + create_endpoint( + sagemaker, + model_name=state.image_model_name, + endpoint_config_name=state.image_endpoint_config_name, + endpoint_name=state.image_endpoint_name, + role_arn=args.role_arn, + image_uri=image_uri, + model_id=IMAGE_MODEL_ID, + instance_type=args.image_instance_type, + startup_timeout_seconds=1800, + ) + wait_for_endpoint(sagemaker, state.image_endpoint_name) + + print(f"Creating video endpoint {state.video_endpoint_name}") + create_endpoint( + sagemaker, + model_name=state.video_model_name, + endpoint_config_name=state.video_endpoint_config_name, + endpoint_name=state.video_endpoint_name, + role_arn=args.role_arn, + image_uri=image_uri, + model_id=VIDEO_MODEL_ID, + instance_type=args.video_instance_type, + startup_timeout_seconds=3600, + async_output_path=f"s3://{bucket}/{state.prefix}/outputs/", + ) + wait_for_endpoint(sagemaker, state.video_endpoint_name) + + print(f"Deployment state written to {args.state_file}") + print(f"Image endpoint: {state.image_endpoint_name}") + print(f"Video endpoint: {state.video_endpoint_name}") + + +if __name__ == "__main__": + main() diff --git a/03-features/vllm-omni-image-video/generate.py b/03-features/vllm-omni-image-video/generate.py new file mode 100644 index 0000000..8539fc5 --- /dev/null +++ b/03-features/vllm-omni-image-video/generate.py @@ -0,0 +1,121 @@ +#!/usr/bin/env python3 +"""Generate an image with FLUX.2-klein and animate it with Wan VACE.""" + +from __future__ import annotations + +import argparse +import time +from datetime import UTC, datetime +from pathlib import Path + +import boto3 + +from vllm_omni_media import ( + DEFAULT_STATE_PATH, + invoke_image, + load_state, + submit_video, + validate_mp4, + wait_for_s3_object, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--image-prompt", + default=( + "Cinematic photograph of a small observatory on a windswept coastal " + "cliff at sunrise, detailed clouds, natural light" + ), + ) + parser.add_argument( + "--video-prompt", + default=( + "Slow camera push-in as clouds move across the sky and grass bends " + "in the wind, stable composition" + ), + ) + parser.add_argument("--image-size", default="1024x1024") + parser.add_argument("--image-steps", type=int, default=4) + parser.add_argument("--width", type=int, default=480) + parser.add_argument("--height", type=int, default=320) + parser.add_argument("--frames", type=int, default=17) + parser.add_argument("--fps", type=int, default=8) + parser.add_argument("--video-steps", type=int, default=4) + parser.add_argument("--guidance-scale", type=float, default=5.0) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--timeout", type=int, default=3600) + parser.add_argument("--output-dir", type=Path, default=Path("outputs")) + parser.add_argument( + "--state-file", + type=Path, + default=DEFAULT_STATE_PATH, + ) + parser.add_argument( + "--keep-request", + action="store_true", + help="Keep the multipart request object in S3 after generation.", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + state = load_state(args.state_file) + session = boto3.Session(region_name=state.region) + runtime = session.client("sagemaker-runtime") + s3 = session.client("s3") + args.output_dir.mkdir(parents=True, exist_ok=True) + stamp = datetime.now(UTC).strftime("%Y%m%d%H%M%S") + + image_started = time.perf_counter() + image_bytes = invoke_image( + runtime, + state, + args.image_prompt, + size=args.image_size, + steps=args.image_steps, + seed=args.seed, + ) + image_seconds = time.perf_counter() - image_started + image_path = args.output_dir / f"{stamp}-flux2-klein.png" + image_path.write_bytes(image_bytes) + print(f"Image: {image_path} ({image_seconds:.1f}s)") + + video_started = time.perf_counter() + output_uri, failure_uri, request_key = submit_video( + runtime, + s3, + state, + args.video_prompt, + image_bytes, + width=args.width, + height=args.height, + num_frames=args.frames, + fps=args.fps, + steps=args.video_steps, + guidance_scale=args.guidance_scale, + seed=args.seed, + ) + try: + video_bytes = wait_for_s3_object( + s3, + output_uri, + failure_uri=failure_uri, + timeout_seconds=args.timeout, + ) + finally: + if not args.keep_request: + s3.delete_object(Bucket=state.bucket, Key=request_key) + + validate_mp4(video_bytes) + video_seconds = time.perf_counter() - video_started + video_path = args.output_dir / f"{stamp}-wan-vace.mp4" + video_path.write_bytes(video_bytes) + print(f"Video: {video_path} ({video_seconds:.1f}s)") + print(f"S3 output: {output_uri}") + + +if __name__ == "__main__": + main() diff --git a/03-features/vllm-omni-image-video/requirements-dev.txt b/03-features/vllm-omni-image-video/requirements-dev.txt new file mode 100644 index 0000000..9c406e1 --- /dev/null +++ b/03-features/vllm-omni-image-video/requirements-dev.txt @@ -0,0 +1,4 @@ +-r requirements.txt +bandit>=1.8,<2 +pytest>=8,<9 +ruff>=0.12,<1 diff --git a/03-features/vllm-omni-image-video/requirements.txt b/03-features/vllm-omni-image-video/requirements.txt new file mode 100644 index 0000000..f9f6605 --- /dev/null +++ b/03-features/vllm-omni-image-video/requirements.txt @@ -0,0 +1,3 @@ +boto3>=1.40,<2 +streamlit>=1.49,<2 +urllib3>=2,<3 diff --git a/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py b/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py new file mode 100644 index 0000000..1dd5319 --- /dev/null +++ b/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py @@ -0,0 +1,256 @@ +import base64 +import json +from pathlib import Path + +import pytest +from botocore.exceptions import ClientError + +from vllm_omni_media import ( + DeploymentState, + build_image_payload, + build_video_multipart, + container_image_uri, + create_endpoint, + decode_image_response, + image_data_url, + load_state, + parse_s3_uri, + save_state, + submit_video, + validate_mp4, + wait_for_s3_object, +) + + +def test_container_image_uri_uses_pinned_dlc_release(): + assert container_image_uri("us-east-1") == ( + "763104351884.dkr.ecr.us-east-1.amazonaws.com/" + "vllm:omni-sagemaker-cuda-v1.6" + ) + + +def test_build_image_payload_uses_flux_klein_defaults(): + payload = build_image_payload("A lighthouse above a stormy sea", seed=7) + + assert payload == { + "model": "black-forest-labs/FLUX.2-klein-4B", + "prompt": "A lighthouse above a stormy sea", + "size": "1024x1024", + "num_inference_steps": 4, + "seed": 7, + } + + +def test_decode_image_response_returns_png_bytes(): + expected = b"\x89PNG\r\n\x1a\nimage" + response = {"data": [{"b64_json": base64.b64encode(expected).decode("ascii")}]} + + assert decode_image_response(json.dumps(response).encode("utf-8")) == expected + + +def test_decode_image_response_rejects_missing_image(): + with pytest.raises(ValueError, match="b64_json"): + decode_image_response(b'{"data": []}') + + +def test_image_data_url_encodes_png(): + assert image_data_url(b"png") == "data:image/png;base64,cG5n" + + +def test_build_video_multipart_contains_reference_and_controls(): + body, content_type = build_video_multipart( + prompt="Slow camera push-in while clouds move", + image_bytes=b"png", + width=480, + height=320, + num_frames=17, + fps=8, + steps=4, + seed=9, + boundary="test-boundary", + ) + + assert content_type == "multipart/form-data; boundary=test-boundary" + text = body.decode("utf-8") + assert 'name="model"\r\n\r\nWan-AI/Wan2.1-VACE-1.3B-diffusers' in text + assert 'name="prompt"\r\n\r\nSlow camera push-in while clouds move' in text + assert 'name="image_reference"\r\n\r\n{"image_url":"data:image/png;base64,cG5n"}' in text + assert 'name="width"\r\n\r\n480' in text + assert 'name="height"\r\n\r\n320' in text + assert 'name="num_frames"\r\n\r\n17' in text + assert 'name="fps"\r\n\r\n8' in text + assert 'name="num_inference_steps"\r\n\r\n4' in text + assert 'name="seed"\r\n\r\n9' in text + assert body.endswith(b"--test-boundary--\r\n") + + +def test_create_video_endpoint_uses_async_inference_and_gpu_ami(): + class FakeSageMaker: + def __init__(self): + self.calls = {} + + def create_model(self, **kwargs): + self.calls["model"] = kwargs + + def create_endpoint_config(self, **kwargs): + self.calls["config"] = kwargs + + def create_endpoint(self, **kwargs): + self.calls["endpoint"] = kwargs + + client = FakeSageMaker() + + create_endpoint( + client, + model_name="video-model", + endpoint_config_name="video-config", + endpoint_name="video-endpoint", + role_arn="arn:aws:iam::111122223333:role/SageMakerRole", + image_uri=container_image_uri("us-east-1"), + model_id="Wan-AI/Wan2.1-VACE-1.3B-diffusers", + instance_type="ml.g6e.xlarge", + startup_timeout_seconds=3600, + async_output_path="s3://example-bucket/outputs/", + ) + + assert client.calls["model"]["PrimaryContainer"]["Environment"] == { + "SM_VLLM_MODEL": "Wan-AI/Wan2.1-VACE-1.3B-diffusers", + "SM_VLLM_VAE_USE_TILING": "true", + } + variant = client.calls["config"]["ProductionVariants"][0] + assert variant["InstanceType"] == "ml.g6e.xlarge" + assert variant["InferenceAmiVersion"] == ( + "al2023-ami-sagemaker-inference-gpu-4-1" + ) + assert client.calls["config"]["AsyncInferenceConfig"] == { + "OutputConfig": {"S3OutputPath": "s3://example-bucket/outputs/"}, + "ClientConfig": {"MaxConcurrentInvocationsPerInstance": 1}, + } + assert client.calls["endpoint"] == { + "EndpointName": "video-endpoint", + "EndpointConfigName": "video-config", + } + + +def test_submit_video_uploads_multipart_request_before_async_invocation(): + class FakeS3: + def __init__(self): + self.request = None + + def put_object(self, **kwargs): + self.request = kwargs + + class FakeRuntime: + def __init__(self): + self.request = None + + def invoke_endpoint_async(self, **kwargs): + self.request = kwargs + return { + "OutputLocation": "s3://example-bucket/outputs/result.out", + "FailureLocation": "s3://example-bucket/failures/result.err", + } + + state = DeploymentState( + region="us-east-1", + bucket="example-bucket", + prefix="vllm-omni-media", + image_model_name="image-model", + image_endpoint_config_name="image-config", + image_endpoint_name="image-endpoint", + video_model_name="video-model", + video_endpoint_config_name="video-config", + video_endpoint_name="video-endpoint", + ) + s3 = FakeS3() + runtime = FakeRuntime() + + output_uri, failure_uri, request_key = submit_video( + runtime, + s3, + state, + "Move the clouds slowly", + b"png", + ) + + assert output_uri == "s3://example-bucket/outputs/result.out" + assert failure_uri == "s3://example-bucket/failures/result.err" + assert request_key.startswith("vllm-omni-media/requests/") + assert s3.request["Bucket"] == "example-bucket" + assert s3.request["Key"] == request_key + assert s3.request["ContentType"].startswith("multipart/form-data; boundary=") + assert runtime.request == { + "EndpointName": "video-endpoint", + "InputLocation": f"s3://example-bucket/{request_key}", + "ContentType": s3.request["ContentType"], + "Accept": "video/mp4", + "CustomAttributes": "route=/v1/videos/sync", + } + + +def test_wait_for_s3_object_surfaces_async_failure(): + class FakeBody: + def __init__(self, value): + self.value = value + + def read(self): + return self.value + + class FakeS3: + def get_object(self, *, Bucket, Key): + if Key == "outputs/result.out": + raise ClientError( + {"Error": {"Code": "NoSuchKey", "Message": "missing"}}, + "GetObject", + ) + return {"Body": FakeBody(b'{"error":"model failed"}')} + + with pytest.raises(RuntimeError, match="model failed"): + wait_for_s3_object( + FakeS3(), + "s3://example-bucket/outputs/result.out", + failure_uri="s3://example-bucket/failures/result.err", + timeout_seconds=1, + poll_seconds=0, + ) + + +def test_state_round_trip(tmp_path: Path): + state = DeploymentState( + region="us-east-1", + bucket="example-bucket", + prefix="vllm-omni-media", + image_model_name="image-model", + image_endpoint_config_name="image-config", + image_endpoint_name="image-endpoint", + video_model_name="video-model", + video_endpoint_config_name="video-config", + video_endpoint_name="video-endpoint", + ) + path = tmp_path / "state.json" + + save_state(state, path) + + assert load_state(path) == state + + +def test_parse_s3_uri(): + assert parse_s3_uri("s3://example-bucket/path/to/output.mp4") == ( + "example-bucket", + "path/to/output.mp4", + ) + + +def test_validate_mp4_accepts_iso_base_media_header(): + validate_mp4(b"\x00\x00\x00\x18ftypisom") + + +def test_validate_mp4_rejects_non_video_response(): + with pytest.raises(ValueError, match="not an MP4"): + validate_mp4(b'{"error":"generation failed"}') + + +@pytest.mark.parametrize("uri", ["https://example.com/file", "s3://bucket", "s3:///key"]) +def test_parse_s3_uri_rejects_invalid_values(uri): + with pytest.raises(ValueError, match="S3 URI"): + parse_s3_uri(uri) diff --git a/03-features/vllm-omni-image-video/vllm_omni_media.py b/03-features/vllm-omni-image-video/vllm_omni_media.py new file mode 100644 index 0000000..58d2a29 --- /dev/null +++ b/03-features/vllm-omni-image-video/vllm_omni_media.py @@ -0,0 +1,476 @@ +"""Shared helpers for the vLLM-Omni image-to-video example.""" + +from __future__ import annotations + +import base64 +import json +import time +import uuid +from dataclasses import asdict, dataclass +from pathlib import Path +from urllib.parse import urlparse + +from botocore.exceptions import ClientError +from urllib3.filepost import encode_multipart_formdata + +DLC_ACCOUNT_ID = "763104351884" +DLC_TAG = "omni-sagemaker-cuda-v1.6" +INFERENCE_AMI_VERSION = "al2023-ami-sagemaker-inference-gpu-4-1" + +IMAGE_MODEL_ID = "black-forest-labs/FLUX.2-klein-4B" +VIDEO_MODEL_ID = "Wan-AI/Wan2.1-VACE-1.3B-diffusers" +IMAGE_ROUTE = "/v1/images/generations" +VIDEO_ROUTE = "/v1/videos/sync" + +DEFAULT_STATE_PATH = Path(__file__).with_name(".vllm_omni_media_state.json") + + +@dataclass(frozen=True) +class DeploymentState: + """Names and locations required to invoke or remove the deployment.""" + + region: str + bucket: str + prefix: str + image_model_name: str + image_endpoint_config_name: str + image_endpoint_name: str + video_model_name: str + video_endpoint_config_name: str + video_endpoint_name: str + + +def container_image_uri(region: str) -> str: + """Return the regional vLLM-Omni Deep Learning Container image URI.""" + + return ( + f"{DLC_ACCOUNT_ID}.dkr.ecr.{region}.amazonaws.com/" + f"vllm:{DLC_TAG}" + ) + + +def build_image_payload( + prompt: str, + *, + size: str = "1024x1024", + steps: int = 4, + seed: int = 42, +) -> dict[str, object]: + """Build a FLUX.2-klein image generation request.""" + + return { + "model": IMAGE_MODEL_ID, + "prompt": prompt, + "size": size, + "num_inference_steps": steps, + "seed": seed, + } + + +def decode_image_response(body: bytes) -> bytes: + """Decode the first base64 image from an OpenAI-compatible response.""" + + try: + response = json.loads(body) + encoded = response["data"][0]["b64_json"] + except (json.JSONDecodeError, KeyError, IndexError, TypeError) as error: + raise ValueError("Image response did not contain data[0].b64_json") from error + + try: + return base64.b64decode(encoded, validate=True) + except (ValueError, TypeError) as error: + raise ValueError("Image response contained invalid base64 data") from error + + +def image_data_url(image_bytes: bytes) -> str: + """Encode PNG bytes as a data URL accepted by vLLM-Omni.""" + + encoded = base64.b64encode(image_bytes).decode("ascii") + return f"data:image/png;base64,{encoded}" + + +def build_video_multipart( + prompt: str, + image_bytes: bytes, + *, + width: int = 480, + height: int = 320, + num_frames: int = 17, + fps: int = 8, + steps: int = 4, + guidance_scale: float = 5.0, + seed: int = 42, + boundary: str | None = None, +) -> tuple[bytes, str]: + """Build the multipart body required by the vLLM-Omni Videos API.""" + + reference = json.dumps( + {"image_url": image_data_url(image_bytes)}, + separators=(",", ":"), + ) + fields = { + "model": VIDEO_MODEL_ID, + "prompt": prompt, + "image_reference": reference, + "width": str(width), + "height": str(height), + "num_frames": str(num_frames), + "fps": str(fps), + "num_inference_steps": str(steps), + "guidance_scale": str(guidance_scale), + "seed": str(seed), + } + return encode_multipart_formdata( + fields, + boundary=boundary or f"vllm-omni-{uuid.uuid4().hex}", + ) + + +def parse_s3_uri(uri: str) -> tuple[str, str]: + """Split a complete S3 URI into bucket and key.""" + + parsed = urlparse(uri) + if parsed.scheme != "s3" or not parsed.netloc or not parsed.path.lstrip("/"): + raise ValueError(f"Invalid S3 URI: {uri}") + return parsed.netloc, parsed.path.lstrip("/") + + +def save_state( + state: DeploymentState, + path: str | Path = DEFAULT_STATE_PATH, +) -> None: + """Persist deployment state for the generation and cleanup scripts.""" + + state_path = Path(path) + state_path.parent.mkdir(parents=True, exist_ok=True) + state_path.write_text( + json.dumps(asdict(state), indent=2) + "\n", + encoding="utf-8", + ) + + +def load_state(path: str | Path = DEFAULT_STATE_PATH) -> DeploymentState: + """Load a previously saved deployment state.""" + + state_path = Path(path) + if not state_path.exists(): + raise FileNotFoundError( + f"Deployment state not found at {state_path}. Run deploy.py first." + ) + return DeploymentState(**json.loads(state_path.read_text(encoding="utf-8"))) + + +def ensure_bucket(s3_client, bucket: str, region: str) -> None: + """Create a private, encrypted S3 bucket when it does not already exist.""" + + try: + s3_client.head_bucket(Bucket=bucket) + return + except ClientError as error: + status = error.response.get("ResponseMetadata", {}).get("HTTPStatusCode") + if status not in {403, 404}: + raise + if status == 403: + raise PermissionError(f"Bucket exists but is not accessible: {bucket}") from error + + create_args: dict[str, object] = {"Bucket": bucket} + if region != "us-east-1": + create_args["CreateBucketConfiguration"] = { + "LocationConstraint": region, + } + s3_client.create_bucket(**create_args) + s3_client.put_public_access_block( + Bucket=bucket, + PublicAccessBlockConfiguration={ + "BlockPublicAcls": True, + "IgnorePublicAcls": True, + "BlockPublicPolicy": True, + "RestrictPublicBuckets": True, + }, + ) + s3_client.put_bucket_encryption( + Bucket=bucket, + ServerSideEncryptionConfiguration={ + "Rules": [ + { + "ApplyServerSideEncryptionByDefault": { + "SSEAlgorithm": "AES256", + } + } + ] + }, + ) + + +def wait_for_endpoint( + sagemaker_client, + endpoint_name: str, + *, + timeout_seconds: int = 3600, + poll_seconds: int = 30, +) -> None: + """Wait until a SageMaker endpoint is in service or fails.""" + + deadline = time.monotonic() + timeout_seconds + previous_status = None + while time.monotonic() < deadline: + description = sagemaker_client.describe_endpoint( + EndpointName=endpoint_name + ) + status = description["EndpointStatus"] + if status != previous_status: + print(f"{endpoint_name}: {status}") + previous_status = status + if status == "InService": + return + if status in {"Failed", "OutOfService"}: + reason = description.get("FailureReason", "No failure reason returned") + raise RuntimeError(f"Endpoint {endpoint_name} failed: {reason}") + time.sleep(poll_seconds) + raise TimeoutError(f"Timed out waiting for endpoint {endpoint_name}") + + +def create_endpoint( + sagemaker_client, + *, + model_name: str, + endpoint_config_name: str, + endpoint_name: str, + role_arn: str, + image_uri: str, + model_id: str, + instance_type: str, + startup_timeout_seconds: int, + async_output_path: str | None = None, +) -> None: + """Create a model, endpoint configuration, and SageMaker endpoint.""" + + environment = {"SM_VLLM_MODEL": model_id} + if model_id == VIDEO_MODEL_ID: + environment["SM_VLLM_VAE_USE_TILING"] = "true" + + sagemaker_client.create_model( + ModelName=model_name, + ExecutionRoleArn=role_arn, + PrimaryContainer={ + "Image": image_uri, + "Environment": environment, + }, + ) + + endpoint_config: dict[str, object] = { + "EndpointConfigName": endpoint_config_name, + "ProductionVariants": [ + { + "VariantName": "AllTraffic", + "ModelName": model_name, + "InstanceType": instance_type, + "InitialInstanceCount": 1, + "InitialVariantWeight": 1.0, + "ContainerStartupHealthCheckTimeoutInSeconds": ( + startup_timeout_seconds + ), + "InferenceAmiVersion": INFERENCE_AMI_VERSION, + } + ], + } + if async_output_path: + endpoint_config["AsyncInferenceConfig"] = { + "OutputConfig": {"S3OutputPath": async_output_path}, + "ClientConfig": {"MaxConcurrentInvocationsPerInstance": 1}, + } + + sagemaker_client.create_endpoint_config(**endpoint_config) + sagemaker_client.create_endpoint( + EndpointName=endpoint_name, + EndpointConfigName=endpoint_config_name, + ) + + +def invoke_image( + runtime_client, + state: DeploymentState, + prompt: str, + *, + size: str = "1024x1024", + steps: int = 4, + seed: int = 42, +) -> bytes: + """Generate an image through the real-time endpoint.""" + + response = runtime_client.invoke_endpoint( + EndpointName=state.image_endpoint_name, + Body=json.dumps( + build_image_payload( + prompt, + size=size, + steps=steps, + seed=seed, + ) + ), + ContentType="application/json", + Accept="application/json", + CustomAttributes=f"route={IMAGE_ROUTE}", + ) + return decode_image_response(response["Body"].read()) + + +def submit_video( + runtime_client, + s3_client, + state: DeploymentState, + prompt: str, + image_bytes: bytes, + *, + width: int = 480, + height: int = 320, + num_frames: int = 17, + fps: int = 8, + steps: int = 4, + guidance_scale: float = 5.0, + seed: int = 42, +) -> tuple[str, str | None, str]: + """Upload a multipart request and submit it to the async video endpoint.""" + + body, content_type = build_video_multipart( + prompt, + image_bytes, + width=width, + height=height, + num_frames=num_frames, + fps=fps, + steps=steps, + guidance_scale=guidance_scale, + seed=seed, + ) + request_key = f"{state.prefix}/requests/{uuid.uuid4().hex}.multipart" + s3_client.put_object( + Bucket=state.bucket, + Key=request_key, + Body=body, + ContentType=content_type, + ServerSideEncryption="AES256", + ) + + response = runtime_client.invoke_endpoint_async( + EndpointName=state.video_endpoint_name, + InputLocation=f"s3://{state.bucket}/{request_key}", + ContentType=content_type, + Accept="video/mp4", + CustomAttributes=f"route={VIDEO_ROUTE}", + ) + return response["OutputLocation"], response.get("FailureLocation"), request_key + + +def wait_for_s3_object( + s3_client, + uri: str, + *, + failure_uri: str | None = None, + timeout_seconds: int = 3600, + poll_seconds: int = 15, +) -> bytes: + """Poll an async inference output location and return its bytes.""" + + bucket, key = parse_s3_uri(uri) + failure_bucket, failure_key = ( + parse_s3_uri(failure_uri) if failure_uri else (None, None) + ) + deadline = time.monotonic() + timeout_seconds + while time.monotonic() < deadline: + try: + response = s3_client.get_object(Bucket=bucket, Key=key) + return response["Body"].read() + except ClientError as error: + code = error.response.get("Error", {}).get("Code") + if code not in {"404", "NoSuchKey", "NotFound"}: + raise + + if failure_bucket and failure_key: + try: + response = s3_client.get_object( + Bucket=failure_bucket, + Key=failure_key, + ) + detail = response["Body"].read().decode("utf-8", errors="replace") + raise RuntimeError(f"Video generation failed: {detail}") + except ClientError as error: + code = error.response.get("Error", {}).get("Code") + if code not in {"404", "NoSuchKey", "NotFound"}: + raise + time.sleep(poll_seconds) + raise TimeoutError(f"Timed out waiting for {uri}") + + +def validate_mp4(video_bytes: bytes) -> None: + """Raise when a response does not have an ISO base media file header.""" + + if len(video_bytes) < 12 or video_bytes[4:8] != b"ftyp": + raise ValueError("Video response is not an MP4 file") + + +def wait_for_endpoint_deleted( + sagemaker_client, + endpoint_name: str, + *, + timeout_seconds: int = 1800, + poll_seconds: int = 15, +) -> None: + """Wait until SageMaker no longer returns an endpoint description.""" + + deadline = time.monotonic() + timeout_seconds + while time.monotonic() < deadline: + try: + sagemaker_client.describe_endpoint(EndpointName=endpoint_name) + except ClientError as error: + code = error.response.get("Error", {}).get("Code") + if code == "ValidationException": + return + raise + time.sleep(poll_seconds) + raise TimeoutError(f"Timed out deleting endpoint {endpoint_name}") + + +def delete_deployment(sagemaker_client, state: DeploymentState) -> None: + """Delete endpoint resources, ignoring resources already removed.""" + + endpoint_names = [ + state.video_endpoint_name, + state.image_endpoint_name, + ] + for endpoint_name in endpoint_names: + try: + sagemaker_client.delete_endpoint(EndpointName=endpoint_name) + except ClientError as error: + code = error.response.get("Error", {}).get("Code") + message = error.response.get("Error", {}).get("Message", "") + if code != "ValidationException" or "Could not find" not in message: + raise + wait_for_endpoint_deleted(sagemaker_client, endpoint_name) + + operations = [ + ( + sagemaker_client.delete_endpoint_config, + {"EndpointConfigName": state.video_endpoint_config_name}, + ), + ( + sagemaker_client.delete_endpoint_config, + {"EndpointConfigName": state.image_endpoint_config_name}, + ), + ( + sagemaker_client.delete_model, + {"ModelName": state.video_model_name}, + ), + ( + sagemaker_client.delete_model, + {"ModelName": state.image_model_name}, + ), + ] + for operation, parameters in operations: + try: + operation(**parameters) + except ClientError as error: + code = error.response.get("Error", {}).get("Code") + message = error.response.get("Error", {}).get("Message", "") + if code != "ValidationException" or "Could not find" not in message: + raise From 7ac8bebcadda994a091a586990e3c023290c1a5f Mon Sep 17 00:00:00 2001 From: Daniel Wirjo Date: Sat, 29 Aug 2026 21:08:48 -0700 Subject: [PATCH 2/3] Make vLLM-Omni deployment resumable --- 03-features/vllm-omni-image-video/README.md | 13 +++ 03-features/vllm-omni-image-video/deploy.py | 79 ++++++++++++------- .../tests/test_vllm_omni_media.py | 64 +++++++++++++++ .../vllm-omni-image-video/vllm_omni_media.py | 63 ++++++++++++--- 4 files changed, 180 insertions(+), 39 deletions(-) diff --git a/03-features/vllm-omni-image-video/README.md b/03-features/vllm-omni-image-video/README.md index 39d5f80..495b0e5 100644 --- a/03-features/vllm-omni-image-video/README.md +++ b/03-features/vllm-omni-image-video/README.md @@ -71,6 +71,19 @@ The deployment script writes resource names to `.vllm_omni_media_state.json`. Keep this file until you finish generation and cleanup. +If your local AWS credentials expire while SageMaker is starting an endpoint, +refresh them and resume from the saved state: + +```bash +python deploy.py \ + --role-arn "$SAGEMAKER_ROLE_ARN" \ + --region us-east-1 \ + --resume +``` + +The resume path reuses resources that already exist and creates only the +remaining resources. + FLUX.2-klein uses a real-time endpoint because its response completes within a single invocation. Wan VACE uses SageMaker Asynchronous Inference because video generation can exceed the response window for real-time inference. The diff --git a/03-features/vllm-omni-image-video/deploy.py b/03-features/vllm-omni-image-video/deploy.py index 88f43a9..4e1f5f0 100644 --- a/03-features/vllm-omni-image-video/deploy.py +++ b/03-features/vllm-omni-image-video/deploy.py @@ -18,6 +18,7 @@ container_image_uri, create_endpoint, ensure_bucket, + load_state, save_state, wait_for_endpoint, ) @@ -44,6 +45,11 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--name-prefix", default="vllm-omni-media") parser.add_argument("--image-instance-type", default="ml.g6.xlarge") parser.add_argument("--video-instance-type", default="ml.g6e.xlarge") + parser.add_argument( + "--resume", + action="store_true", + help="Resume the deployment recorded in --state-file.", + ) parser.add_argument( "--state-file", type=Path, @@ -55,37 +61,56 @@ def parse_args() -> argparse.Namespace: def main() -> None: args = parse_args() session = boto3.Session(region_name=args.region) - sts = session.client("sts") - account_id = sts.get_caller_identity()["Account"] - bucket = args.bucket or f"sagemaker-{args.region}-{account_id}" sagemaker = session.client("sagemaker") s3 = session.client("s3") - ensure_bucket(s3, bucket, args.region) - timestamp = datetime.now(UTC).strftime("%Y%m%d%H%M%S") - state = DeploymentState( - region=args.region, - bucket=bucket, - prefix=args.prefix.strip("/"), - image_model_name=resource_name(args.name_prefix, "image-model", timestamp), - image_endpoint_config_name=resource_name( - args.name_prefix, "image-config", timestamp - ), - image_endpoint_name=resource_name( - args.name_prefix, "image-endpoint", timestamp - ), - video_model_name=resource_name(args.name_prefix, "video-model", timestamp), - video_endpoint_config_name=resource_name( - args.name_prefix, "video-config", timestamp - ), - video_endpoint_name=resource_name( - args.name_prefix, "video-endpoint", timestamp - ), - ) - save_state(state, args.state_file) + if args.resume: + state = load_state(args.state_file) + if state.region != args.region: + raise ValueError( + f"State file uses {state.region}, but --region is {args.region}." + ) + bucket = state.bucket + print(f"Resuming deployment from {args.state_file}") + else: + if args.state_file.exists(): + raise FileExistsError( + f"State file already exists at {args.state_file}. " + "Use --resume or remove it after cleanup." + ) + sts = session.client("sts") + account_id = sts.get_caller_identity()["Account"] + bucket = args.bucket or f"sagemaker-{args.region}-{account_id}" + timestamp = datetime.now(UTC).strftime("%Y%m%d%H%M%S") + state = DeploymentState( + region=args.region, + bucket=bucket, + prefix=args.prefix.strip("/"), + image_model_name=resource_name( + args.name_prefix, "image-model", timestamp + ), + image_endpoint_config_name=resource_name( + args.name_prefix, "image-config", timestamp + ), + image_endpoint_name=resource_name( + args.name_prefix, "image-endpoint", timestamp + ), + video_model_name=resource_name( + args.name_prefix, "video-model", timestamp + ), + video_endpoint_config_name=resource_name( + args.name_prefix, "video-config", timestamp + ), + video_endpoint_name=resource_name( + args.name_prefix, "video-endpoint", timestamp + ), + ) + save_state(state, args.state_file) + + ensure_bucket(s3, bucket, args.region) image_uri = container_image_uri(args.region) - print(f"Creating image endpoint {state.image_endpoint_name}") + print(f"Preparing image endpoint {state.image_endpoint_name}") create_endpoint( sagemaker, model_name=state.image_model_name, @@ -99,7 +124,7 @@ def main() -> None: ) wait_for_endpoint(sagemaker, state.image_endpoint_name) - print(f"Creating video endpoint {state.video_endpoint_name}") + print(f"Preparing video endpoint {state.video_endpoint_name}") create_endpoint( sagemaker, model_name=state.video_model_name, diff --git a/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py b/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py index 1dd5319..08fd4d9 100644 --- a/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py +++ b/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py @@ -89,6 +89,27 @@ class FakeSageMaker: def __init__(self): self.calls = {} + @staticmethod + def _missing(operation): + raise ClientError( + { + "Error": { + "Code": "ValidationException", + "Message": "resource not found", + } + }, + operation, + ) + + def describe_model(self, **kwargs): + self._missing("DescribeModel") + + def describe_endpoint_config(self, **kwargs): + self._missing("DescribeEndpointConfig") + + def describe_endpoint(self, **kwargs): + self._missing("DescribeEndpoint") + def create_model(self, **kwargs): self.calls["model"] = kwargs @@ -132,6 +153,49 @@ def create_endpoint(self, **kwargs): } +def test_create_endpoint_reuses_existing_resources(): + class FakeSageMaker: + def __init__(self): + self.create_calls = [] + + def describe_model(self, **kwargs): + return {"ModelName": kwargs["ModelName"]} + + def describe_endpoint_config(self, **kwargs): + return {"EndpointConfigName": kwargs["EndpointConfigName"]} + + def describe_endpoint(self, **kwargs): + return { + "EndpointName": kwargs["EndpointName"], + "EndpointStatus": "InService", + } + + def create_model(self, **kwargs): + self.create_calls.append(("model", kwargs)) + + def create_endpoint_config(self, **kwargs): + self.create_calls.append(("config", kwargs)) + + def create_endpoint(self, **kwargs): + self.create_calls.append(("endpoint", kwargs)) + + client = FakeSageMaker() + + create_endpoint( + client, + model_name="image-model", + endpoint_config_name="image-config", + endpoint_name="image-endpoint", + role_arn="arn:aws:iam::111122223333:role/SageMakerRole", + image_uri=container_image_uri("us-east-1"), + model_id="black-forest-labs/FLUX.2-klein-4B", + instance_type="ml.g6.xlarge", + startup_timeout_seconds=1800, + ) + + assert client.create_calls == [] + + def test_submit_video_uploads_multipart_request_before_async_invocation(): class FakeS3: def __init__(self): diff --git a/03-features/vllm-omni-image-video/vllm_omni_media.py b/03-features/vllm-omni-image-video/vllm_omni_media.py index 58d2a29..5f7733e 100644 --- a/03-features/vllm-omni-image-video/vllm_omni_media.py +++ b/03-features/vllm-omni-image-video/vllm_omni_media.py @@ -230,6 +230,23 @@ def wait_for_endpoint( raise TimeoutError(f"Timed out waiting for endpoint {endpoint_name}") +def sagemaker_resource_exists( + sagemaker_client, + describe_method: str, + **kwargs, +) -> bool: + """Return whether a named SageMaker resource exists.""" + + try: + getattr(sagemaker_client, describe_method)(**kwargs) + return True + except ClientError as error: + code = error.response.get("Error", {}).get("Code") + if code == "ValidationException": + return False + raise + + def create_endpoint( sagemaker_client, *, @@ -243,20 +260,27 @@ def create_endpoint( startup_timeout_seconds: int, async_output_path: str | None = None, ) -> None: - """Create a model, endpoint configuration, and SageMaker endpoint.""" + """Create missing resources for a SageMaker endpoint.""" environment = {"SM_VLLM_MODEL": model_id} if model_id == VIDEO_MODEL_ID: environment["SM_VLLM_VAE_USE_TILING"] = "true" - sagemaker_client.create_model( + if sagemaker_resource_exists( + sagemaker_client, + "describe_model", ModelName=model_name, - ExecutionRoleArn=role_arn, - PrimaryContainer={ - "Image": image_uri, - "Environment": environment, - }, - ) + ): + print(f"Reusing model {model_name}") + else: + sagemaker_client.create_model( + ModelName=model_name, + ExecutionRoleArn=role_arn, + PrimaryContainer={ + "Image": image_uri, + "Environment": environment, + }, + ) endpoint_config: dict[str, object] = { "EndpointConfigName": endpoint_config_name, @@ -280,11 +304,26 @@ def create_endpoint( "ClientConfig": {"MaxConcurrentInvocationsPerInstance": 1}, } - sagemaker_client.create_endpoint_config(**endpoint_config) - sagemaker_client.create_endpoint( - EndpointName=endpoint_name, + if sagemaker_resource_exists( + sagemaker_client, + "describe_endpoint_config", EndpointConfigName=endpoint_config_name, - ) + ): + print(f"Reusing endpoint configuration {endpoint_config_name}") + else: + sagemaker_client.create_endpoint_config(**endpoint_config) + + if sagemaker_resource_exists( + sagemaker_client, + "describe_endpoint", + EndpointName=endpoint_name, + ): + print(f"Reusing endpoint {endpoint_name}") + else: + sagemaker_client.create_endpoint( + EndpointName=endpoint_name, + EndpointConfigName=endpoint_config_name, + ) def invoke_image( From 7aa491d834f91fb2a69c821281c3774428682e18 Mon Sep 17 00:00:00 2001 From: Daniel Wirjo Date: Sat, 29 Aug 2026 22:20:03 -0700 Subject: [PATCH 3/3] Harden and validate vLLM-Omni media workflow --- 03-features/vllm-omni-image-video/README.md | 22 ++++--- 03-features/vllm-omni-image-video/app.py | 17 ++--- 03-features/vllm-omni-image-video/deploy.py | 1 + 03-features/vllm-omni-image-video/generate.py | 7 +- .../vllm-omni-image-video/requirements.txt | 1 + .../tests/test_vllm_omni_media.py | 33 ++++++++-- .../vllm-omni-image-video/vllm_omni_media.py | 65 ++++++++++++++++--- 7 files changed, 112 insertions(+), 34 deletions(-) diff --git a/03-features/vllm-omni-image-video/README.md b/03-features/vllm-omni-image-video/README.md index 495b0e5..55d9712 100644 --- a/03-features/vllm-omni-image-video/README.md +++ b/03-features/vllm-omni-image-video/README.md @@ -25,10 +25,11 @@ flowchart LR ``` The image endpoint returns an OpenAI-compatible JSON response containing a -base64-encoded PNG. The sample places that PNG in the video request as a base64 -data URL. It uploads the complete multipart request to Amazon S3, invokes -`/v1/videos/sync` through the asynchronous endpoint, and downloads the MP4 from -the returned output location. +base64-encoded PNG. The sample resizes that image to the video dimensions and +encodes it as a compact JPEG data URL so the multipart field remains below the +server's per-part size. It uploads the complete multipart request to Amazon S3, +invokes `/v1/videos/sync` through the asynchronous endpoint, and downloads the +MP4 from the returned output location. ## Prerequisites @@ -96,16 +97,16 @@ Run the complete workflow: ```bash python generate.py \ --image-prompt "Cinematic photograph of a coastal observatory at sunrise" \ - --video-prompt "Slow camera push-in as clouds move and grass bends in the wind" + --video-prompt "Slow camera push-in toward the coastal observatory; preserve the building and coastline" ``` The script writes the PNG and MP4 to `outputs/`. It also prints the Amazon S3 output location returned by SageMaker Asynchronous Inference. -The default video settings use 17 frames and four diffusion steps to keep a -validation run short. Increase `--frames` and `--video-steps` after you confirm -the deployment, then assess output quality and instance memory for your -workload. +The default video settings use 17 frames and 30 diffusion steps. This matches +the step count in the vLLM-Omni Wan VACE recipe while keeping the clip short. +Use `--video-steps 4` only for a quick endpoint smoke test, then assess output +quality, latency, and instance memory with your production settings. ## Run the Streamlit application @@ -132,7 +133,8 @@ OpenAI-compatible route: The Videos API accepts multipart form data. This sample pre-builds the multipart body before uploading it to Amazon S3, which keeps the request format explicit and works across vLLM-Omni DLC releases that accept multipart -requests. +requests. The asynchronous endpoint writes successful responses to `outputs/` +and invocation errors to `failures/` under the sample Amazon S3 prefix. ## Clean up diff --git a/03-features/vllm-omni-image-video/app.py b/03-features/vllm-omni-image-video/app.py index 7e94fa4..450fdc5 100644 --- a/03-features/vllm-omni-image-video/app.py +++ b/03-features/vllm-omni-image-video/app.py @@ -47,7 +47,7 @@ with seed_col: seed = st.number_input("Seed", min_value=0, value=42, step=1) -if st.button("Generate image", type="primary", use_container_width=True): +if st.button("Generate image", type="primary", width="stretch"): with st.spinner("Generating image"): started = time.perf_counter() st.session_state.image_bytes = invoke_image( @@ -64,22 +64,23 @@ st.image( st.session_state.image_bytes, caption=f"FLUX.2-klein output in {st.session_state.image_seconds:.1f}s", - use_container_width=True, + width="stretch", ) st.download_button( "Download image", st.session_state.image_bytes, "flux2-klein.png", "image/png", - use_container_width=True, + width="stretch", ) st.divider() video_prompt = st.text_area( "Motion prompt", ( - "Slow camera push-in as clouds move across the sky and grass bends " - "in the wind, stable composition" + "Slow camera push-in toward the coastal observatory as clouds drift " + "across the sky and ocean waves move below; preserve the building, " + "coastline, and composition" ), ) frames_col, fps_col, steps_col = st.columns(3) @@ -88,12 +89,12 @@ with fps_col: fps = st.select_slider("Frames per second", options=[4, 8, 16], value=8) with steps_col: - steps = st.select_slider("Inference steps", options=[4, 8, 16, 30], value=4) + steps = st.select_slider("Inference steps", options=[4, 8, 16, 30], value=30) if st.button( "Generate video", disabled="image_bytes" not in st.session_state, - use_container_width=True, + width="stretch", ): with st.status("Generating video", expanded=True) as status: status.write("Uploading the multipart request to Amazon S3") @@ -134,5 +135,5 @@ st.session_state.video_bytes, "wan-vace.mp4", "video/mp4", - use_container_width=True, + width="stretch", ) diff --git a/03-features/vllm-omni-image-video/deploy.py b/03-features/vllm-omni-image-video/deploy.py index 4e1f5f0..6ef72a7 100644 --- a/03-features/vllm-omni-image-video/deploy.py +++ b/03-features/vllm-omni-image-video/deploy.py @@ -136,6 +136,7 @@ def main() -> None: instance_type=args.video_instance_type, startup_timeout_seconds=3600, async_output_path=f"s3://{bucket}/{state.prefix}/outputs/", + async_failure_path=f"s3://{bucket}/{state.prefix}/failures/", ) wait_for_endpoint(sagemaker, state.video_endpoint_name) diff --git a/03-features/vllm-omni-image-video/generate.py b/03-features/vllm-omni-image-video/generate.py index 8539fc5..708192e 100644 --- a/03-features/vllm-omni-image-video/generate.py +++ b/03-features/vllm-omni-image-video/generate.py @@ -32,8 +32,9 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--video-prompt", default=( - "Slow camera push-in as clouds move across the sky and grass bends " - "in the wind, stable composition" + "Slow camera push-in toward the coastal observatory as clouds drift " + "across the sky and ocean waves move below; preserve the building, " + "coastline, and composition" ), ) parser.add_argument("--image-size", default="1024x1024") @@ -42,7 +43,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--height", type=int, default=320) parser.add_argument("--frames", type=int, default=17) parser.add_argument("--fps", type=int, default=8) - parser.add_argument("--video-steps", type=int, default=4) + parser.add_argument("--video-steps", type=int, default=30) parser.add_argument("--guidance-scale", type=float, default=5.0) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--timeout", type=int, default=3600) diff --git a/03-features/vllm-omni-image-video/requirements.txt b/03-features/vllm-omni-image-video/requirements.txt index f9f6605..435d65c 100644 --- a/03-features/vllm-omni-image-video/requirements.txt +++ b/03-features/vllm-omni-image-video/requirements.txt @@ -1,3 +1,4 @@ boto3>=1.40,<2 +pillow>=11,<13 streamlit>=1.49,<2 urllib3>=2,<3 diff --git a/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py b/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py index 08fd4d9..7d9d6ee 100644 --- a/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py +++ b/03-features/vllm-omni-image-video/tests/test_vllm_omni_media.py @@ -1,9 +1,11 @@ import base64 import json +from io import BytesIO from pathlib import Path import pytest from botocore.exceptions import ClientError +from PIL import Image from vllm_omni_media import ( DeploymentState, @@ -15,12 +17,18 @@ image_data_url, load_state, parse_s3_uri, + prepare_video_reference, save_state, submit_video, validate_mp4, wait_for_s3_object, ) +def make_png(size: tuple[int, int] = (1, 1)) -> bytes: + output = BytesIO() + Image.new("RGB", size, "navy").save(output, format="PNG") + return output.getvalue() + def test_container_image_uri_uses_pinned_dlc_release(): assert container_image_uri("us-east-1") == ( @@ -65,7 +73,6 @@ def test_build_video_multipart_contains_reference_and_controls(): height=320, num_frames=17, fps=8, - steps=4, seed=9, boundary="test-boundary", ) @@ -79,7 +86,7 @@ def test_build_video_multipart_contains_reference_and_controls(): assert 'name="height"\r\n\r\n320' in text assert 'name="num_frames"\r\n\r\n17' in text assert 'name="fps"\r\n\r\n8' in text - assert 'name="num_inference_steps"\r\n\r\n4' in text + assert 'name="num_inference_steps"\r\n\r\n30' in text assert 'name="seed"\r\n\r\n9' in text assert body.endswith(b"--test-boundary--\r\n") @@ -132,6 +139,7 @@ def create_endpoint(self, **kwargs): instance_type="ml.g6e.xlarge", startup_timeout_seconds=3600, async_output_path="s3://example-bucket/outputs/", + async_failure_path="s3://example-bucket/failures/", ) assert client.calls["model"]["PrimaryContainer"]["Environment"] == { @@ -144,7 +152,10 @@ def create_endpoint(self, **kwargs): "al2023-ami-sagemaker-inference-gpu-4-1" ) assert client.calls["config"]["AsyncInferenceConfig"] == { - "OutputConfig": {"S3OutputPath": "s3://example-bucket/outputs/"}, + "OutputConfig": { + "S3OutputPath": "s3://example-bucket/outputs/", + "S3FailurePath": "s3://example-bucket/failures/", + }, "ClientConfig": {"MaxConcurrentInvocationsPerInstance": 1}, } assert client.calls["endpoint"] == { @@ -196,6 +207,20 @@ def create_endpoint(self, **kwargs): assert client.create_calls == [] +def test_prepare_video_reference_resizes_to_compact_jpeg(): + reference = prepare_video_reference( + make_png((1024, 1024)), + width=480, + height=320, + ) + + assert reference.startswith(b"\xff\xd8") + assert len(reference) < 1_000_000 + with Image.open(BytesIO(reference)) as image: + assert image.size == (480, 320) + assert image.format == "JPEG" + + def test_submit_video_uploads_multipart_request_before_async_invocation(): class FakeS3: def __init__(self): @@ -234,7 +259,7 @@ def invoke_endpoint_async(self, **kwargs): s3, state, "Move the clouds slowly", - b"png", + make_png(), ) assert output_uri == "s3://example-bucket/outputs/result.out" diff --git a/03-features/vllm-omni-image-video/vllm_omni_media.py b/03-features/vllm-omni-image-video/vllm_omni_media.py index 5f7733e..7df4a95 100644 --- a/03-features/vllm-omni-image-video/vllm_omni_media.py +++ b/03-features/vllm-omni-image-video/vllm_omni_media.py @@ -7,10 +7,12 @@ import time import uuid from dataclasses import asdict, dataclass +from io import BytesIO from pathlib import Path from urllib.parse import urlparse from botocore.exceptions import ClientError +from PIL import Image, ImageOps, UnidentifiedImageError from urllib3.filepost import encode_multipart_formdata DLC_ACCOUNT_ID = "763104351884" @@ -82,11 +84,45 @@ def decode_image_response(body: bytes) -> bytes: raise ValueError("Image response contained invalid base64 data") from error -def image_data_url(image_bytes: bytes) -> str: - """Encode PNG bytes as a data URL accepted by vLLM-Omni.""" +def image_data_url( + image_bytes: bytes, + media_type: str = "image/png", +) -> str: + """Encode image bytes as a data URL accepted by vLLM-Omni.""" encoded = base64.b64encode(image_bytes).decode("ascii") - return f"data:image/png;base64,{encoded}" + return f"data:{media_type};base64,{encoded}" + + +def prepare_video_reference( + image_bytes: bytes, + *, + width: int, + height: int, +) -> bytes: + """Resize the source image to a compact JPEG for the multipart request.""" + + if width <= 0 or height <= 0: + raise ValueError("Video width and height must be positive") + + try: + with Image.open(BytesIO(image_bytes)) as source: + reference = ImageOps.fit( + source.convert("RGB"), + (width, height), + method=Image.Resampling.LANCZOS, + ) + except (OSError, UnidentifiedImageError) as error: + raise ValueError("Source image could not be decoded") from error + + output = BytesIO() + reference.save( + output, + format="JPEG", + quality=90, + optimize=True, + ) + return output.getvalue() def build_video_multipart( @@ -97,15 +133,16 @@ def build_video_multipart( height: int = 320, num_frames: int = 17, fps: int = 8, - steps: int = 4, + steps: int = 30, guidance_scale: float = 5.0, seed: int = 42, + image_media_type: str = "image/png", boundary: str | None = None, ) -> tuple[bytes, str]: """Build the multipart body required by the vLLM-Omni Videos API.""" reference = json.dumps( - {"image_url": image_data_url(image_bytes)}, + {"image_url": image_data_url(image_bytes, image_media_type)}, separators=(",", ":"), ) fields = { @@ -259,6 +296,7 @@ def create_endpoint( instance_type: str, startup_timeout_seconds: int, async_output_path: str | None = None, + async_failure_path: str | None = None, ) -> None: """Create missing resources for a SageMaker endpoint.""" @@ -299,8 +337,11 @@ def create_endpoint( ], } if async_output_path: + output_config = {"S3OutputPath": async_output_path} + if async_failure_path: + output_config["S3FailurePath"] = async_failure_path endpoint_config["AsyncInferenceConfig"] = { - "OutputConfig": {"S3OutputPath": async_output_path}, + "OutputConfig": output_config, "ClientConfig": {"MaxConcurrentInvocationsPerInstance": 1}, } @@ -365,15 +406,20 @@ def submit_video( height: int = 320, num_frames: int = 17, fps: int = 8, - steps: int = 4, + steps: int = 30, guidance_scale: float = 5.0, seed: int = 42, ) -> tuple[str, str | None, str]: """Upload a multipart request and submit it to the async video endpoint.""" + reference_image = prepare_video_reference( + image_bytes, + width=width, + height=height, + ) body, content_type = build_video_multipart( prompt, - image_bytes, + reference_image, width=width, height=height, num_frames=num_frames, @@ -381,6 +427,7 @@ def submit_video( steps=steps, guidance_scale=guidance_scale, seed=seed, + image_media_type="image/jpeg", ) request_key = f"{state.prefix}/requests/{uuid.uuid4().hex}.multipart" s3_client.put_object( @@ -407,7 +454,7 @@ def wait_for_s3_object( *, failure_uri: str | None = None, timeout_seconds: int = 3600, - poll_seconds: int = 15, + poll_seconds: int = 5, ) -> bytes: """Poll an async inference output location and return its bytes."""