Skip to content
Closed
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
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
#!/bin/bash
# SageMaker entry script - installs uv and runs sm_job_runner.py with dependencies

set -e

echo "Installing uv..."
curl -LsSf https://astral.sh/uv/install.sh | sh
export PATH="$HOME/.local/bin:$PATH"

echo "Running pipeline with uv..."
cd /opt/ml/input/data/code
uv run sm_job_runner.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
#!/usr/bin/env python3
# /// script
# requires-python = ">=3.11"
# dependencies = [
# "boto3",
# "s3fs",
# "torch",
# "datasets[s3]>=4.0.0",
# "pyarrow>=12.0.0",
# "numpy",
# "pillow",
# "requests",
# "openai",
# "huggingface-hub[hf_transfer]",
# "rich",
# ]
# ///
"""
SageMaker Training Job entry point for the DeepSeek OCR pipeline.

This script is run via entry.sh which installs uv and executes: uv run sm_job_runner.py
Code is automatically available at /opt/ml/input/data/code via SageMaker SourceCode.

Environment Variables:
PIPELINE_STAGE: Stage to run (extract, describe, assemble)

SageMaker Environment Variables (automatically set):
SM_MODEL_DIR: /opt/ml/model
SM_OUTPUT_DATA_DIR: /opt/ml/output/data
"""
from __future__ import annotations

import json
import logging
import os
import sys
from pathlib import Path


# Code is at /opt/ml/input/data/code (SageMaker SourceCode)
CODE_DIR = Path("/opt/ml/input/data/code")


def setup_logging():
"""Configure logging."""
level = os.environ.get("LOG_LEVEL", "INFO").upper()
logging.basicConfig(
level=level,
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
force=True,
)


def load_hyperparameters() -> dict:
"""Load hyperparameters from SageMaker.

SageMaker passes hyperparameters as a JSON file at /opt/ml/input/config/hyperparameters.json
All values are strings, so we convert them to environment variables.
"""
hp_file = Path("/opt/ml/input/config/hyperparameters.json")
if not hp_file.exists():
logging.info("No hyperparameters file found at %s", hp_file)
return {}

with hp_file.open() as f:
hyperparameters = json.load(f)

logging.info("Loaded hyperparameters: %s", list(hyperparameters.keys()))

# Set hyperparameters as environment variables
for key, value in hyperparameters.items():
# SageMaker wraps values in quotes, strip them
if isinstance(value, str):
value = value.strip('"').strip("'")
os.environ[key] = str(value)

return hyperparameters


def write_success_marker():
"""Write a success marker file for SageMaker."""
model_dir = Path(os.environ.get("SM_MODEL_DIR", "/opt/ml/model"))
model_dir.mkdir(parents=True, exist_ok=True)

success_file = model_dir / "_SUCCESS"
success_file.write_text("Pipeline completed successfully\n")
logging.info("Wrote success marker to %s", success_file)


def main() -> None:
"""Main entry point for SageMaker training job."""
setup_logging()
logger = logging.getLogger(__name__)

logger.info("Starting SageMaker OCR pipeline job")
logger.info("Python version: %s", sys.version)
logger.info("Working directory: %s", os.getcwd())

# Log SageMaker environment
sm_vars = {k: v for k, v in os.environ.items() if k.startswith("SM_")}
logger.info("SageMaker environment: %s", json.dumps(sm_vars, indent=2))

# Load hyperparameters from SageMaker
load_hyperparameters()

# Add code directory to path
sys.path.insert(0, str(CODE_DIR))
logger.info("Code directory: %s", CODE_DIR)

# Import and run pipeline
try:
from llm_ocr.cli import main as pipeline_main
pipeline_main()
write_success_marker()
logger.info("Pipeline completed successfully")
except Exception as exc:
logger.exception("Pipeline failed: %s", exc)
# Write failure info for debugging
failure_file = Path("/opt/ml/output/failure")
failure_file.parent.mkdir(parents=True, exist_ok=True)
failure_file.write_text(str(exc))
raise


if __name__ == "__main__":
main()
39 changes: 39 additions & 0 deletions docs/sagemaker/notebooks/sagemaker-sdk/llm_ocr/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
"""DeepSeek OCR pipeline package.

This package provides tools for running batch OCR inference with DeepSeek-OCR
and similar vision-language models across different cloud platforms.

Example usage:
from llm_ocr import DeepSeekClient, ExtractSettings, run_stage_extract

client = DeepSeekClient(base_url="http://localhost:8000/v1")
settings = ExtractSettings.from_env(client)
run_stage_extract(settings)
"""

# Stage runners - main entry points for pipeline stages
from llm_ocr.stages import (
run_stage_assemble as run_stage_assemble,
run_stage_describe as run_stage_describe,
run_stage_extract as run_stage_extract,
)

# Configuration classes
from llm_ocr.config import (
AssembleSettings as AssembleSettings,
DescribeSettings as DescribeSettings,
ExtractSettings as ExtractSettings,
InferenceSettings as InferenceSettings,
)

# Inference client
from llm_ocr.server import DeepSeekClient as DeepSeekClient

# Storage backends
from llm_ocr.storage import (
DatasetStorage as DatasetStorage,
GCSStorage as GCSStorage,
HFHubStorage as HFHubStorage,
S3Storage as S3Storage,
get_storage as get_storage,
)
6 changes: 6 additions & 0 deletions docs/sagemaker/notebooks/sagemaker-sdk/llm_ocr/__main__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
"""Entry point for `python -m llm_ocr`."""

from llm_ocr.cli import main

if __name__ == "__main__":
main()
111 changes: 111 additions & 0 deletions docs/sagemaker/notebooks/sagemaker-sdk/llm_ocr/cli.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,111 @@
"""CLI entrypoint for the DeepSeek OCR pipeline."""

from __future__ import annotations

import logging

from .config import AssembleSettings, DescribeSettings, ExtractSettings, env
from .server import (
DeepSeekClient,
base_url_from_env,
launch_vllm,
should_launch_server,
shutdown_server,
wait_for_server,
)
from .stages import run_stage_assemble, run_stage_describe, run_stage_extract

LOGGER = logging.getLogger(__name__)


def _setup_logging() -> None:
"""Configure logging with optional rich handler."""
level = env("LOG_LEVEL", "INFO").upper()
try:
from rich.console import Console
from rich.logging import RichHandler

console = Console(
force_terminal=env("FORCE_COLOR", "").lower() in {"1", "true"}
)
handler = RichHandler(
console=console, show_time=True, show_level=True, rich_tracebacks=True
)
logging.basicConfig(
level=level, format="%(message)s", handlers=[handler], force=True
)
except ImportError:
logging.basicConfig(
level=level,
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
force=True,
)


def _create_client(
max_tokens: int, temperature: float, inference_settings
) -> DeepSeekClient:
"""Create DeepSeek client from environment."""
return DeepSeekClient(
base_url=base_url_from_env(),
model_name=env("SERVED_MODEL_NAME", "deepseek-ocr"),
max_tokens=max_tokens,
temperature=temperature,
request_timeout=inference_settings.request_timeout,
max_retries=inference_settings.max_retries,
retry_backoff_seconds=inference_settings.retry_backoff,
)


def main() -> None:
"""Main entry point for the pipeline CLI."""
_setup_logging()

stage = env("PIPELINE_STAGE", "extract").lower()
if stage not in {"extract", "describe", "assemble"}:
raise ValueError(f"Unsupported stage: {stage}")

needs_server = stage in {"extract", "describe"}
launch_server = should_launch_server() and needs_server
server_process = None

try:
if launch_server:
server_process = launch_vllm()

if needs_server:
base_url = base_url_from_env()
health_url = env("HEALTH_URL", f"{base_url}/health")
LOGGER.info("Waiting for server at %s", health_url)
if not wait_for_server(health_url):
raise RuntimeError("vLLM server did not become ready in time")

if stage == "extract":
from .config import InferenceSettings

inference = InferenceSettings.from_env("EXTRACT")
max_tokens = env("DOC_MAX_TOKENS", 2048, int)
temperature = env("DOC_TEMPERATURE", 0.0, float)
client = _create_client(max_tokens, temperature, inference)
settings = ExtractSettings.from_env(client)
settings.inference = inference
run_stage_extract(settings)

elif stage == "describe":
from .config import InferenceSettings

inference = InferenceSettings.from_env("DESCRIBE")
max_tokens = env("FIGURE_MAX_TOKENS", 512, int)
temperature = env("FIGURE_TEMPERATURE", 0.0, float)
client = _create_client(max_tokens, temperature, inference)
settings = DescribeSettings.from_env(client)
settings.inference = inference
run_stage_describe(settings)

elif stage == "assemble":
settings = AssembleSettings.from_env()
run_stage_assemble(settings)

finally:
if server_process is not None:
shutdown_server(server_process)
Loading