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
36 changes: 34 additions & 2 deletions .github/workflows/publish_translation_worker.yml
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ jobs:
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}

publish-inference-worker:
publish-inference-c2translate-worker:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
Expand All @@ -101,7 +101,39 @@ jobs:
- name: Build and push image
uses: docker/build-push-action@v7
with:
target: inference-worker
target: inference-c2translate-worker
context: ./workers/translation-worker
platforms: linux/amd64
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}

publish-inference-torch-worker:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- name: Docker meta
id: meta
uses: docker/metadata-action@v6
with:
images: icij/datashare-translation-inference-worker
tags: |
type=match,pattern=translation-worker-(.*),group=1

- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v4

- name: Login to Docker Hub
uses: docker/login-action@v4
with:
# You'll need to set these secrets
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_TOKEN }}

- name: Build and push image
uses: docker/build-push-action@v7
with:
target: inference-torch-worker
context: ./workers/translation-worker
platforms: linux/amd64
push: true
Expand Down
25 changes: 14 additions & 11 deletions workers/extract-worker/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,7 @@ RUN apt update && \
rm -rf /var/lib/apt/lists/*


FROM extract-base-builder AS extract-cpu-build
FROM extract-base-builder AS extract-cpu-builder
# Install deps first to optimize layer cache
RUN --mount=type=cache,target=/root/.cache/uv \
--mount=type=bind,source=uv.dist.lock,target=uv.lock \
Expand All @@ -91,11 +91,12 @@ RUN rm -rf ~/.cache/pip

FROM runtime-base AS extract-cpu-worker

COPY --from=extract-cpu-build /app /app
COPY --from=extract-cpu-builder /app /app
COPY --from=extract-cpu-builder /root/.local/share/uv/python /root/.local/share/uv/python
ENTRYPOINT ["entrypoints/extract_cpu_worker.sh"]


FROM extract-base-builder AS extract-gpu-build
FROM extract-base-builder AS extract-gpu-builder
## Copy the flash-attn wheel
#COPY --from=flash-attn-builder /flash-attn-wheel /flash-attn-wheel
## Install the wheel
Expand All @@ -122,12 +123,12 @@ FROM runtime-base AS extract-gpu-worker
RUN apt update && \
apt install -y --no-install-recommends tesseract-ocr && \
rm -rf /var/lib/apt/lists/*
COPY --from=extract-gpu-build /app /app

COPY --from=extract-gpu-builder /app /app
COPY --from=extract-gpu-builder /root/.local/share/uv/python /root/.local/share/uv/python
ENTRYPOINT ["entrypoints/extract_gpu_worker.sh"]


FROM runtime-base AS extract-cpu-mineru-build
FROM runtime-base AS extract-cpu-mineru-builder
# Install deps first to optimize layer cache
RUN --mount=type=cache,target=/root/.cache/uv \
--mount=type=bind,source=uv.dist.lock,target=uv.lock \
Expand All @@ -138,7 +139,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
ADD uv.dist.lock ./uv.lock
ADD pyproject.toml README.md ./
ADD extract_worker ./extract_worker/
ADD entrypoints/extract_cpu_worker.sh ./entrypoints/extract_cpu_worker.sh
ADD entrypoints/extract_cpu_mineru_worker.sh ./entrypoints/extract_cpu_mineru_worker.sh

# Then install service
RUN --mount=type=cache,target=/root/.cache/uv uv sync -v --frozen --no-editable --extra mineru --extra cpu
Expand All @@ -147,11 +148,12 @@ RUN rm -rf ~/.cache/pip
# Slim runtime (as extract-cpu-worker): no toolchain, keep tesseract
FROM runtime-base AS extract-cpu-mineru-worker

COPY --from=extract-cpu-mineru-build /app /app
COPY --from=extract-cpu-mineru-builder /app /app
COPY --from=extract-cpu-mineru-builder /root/.local/share/uv/python /root/.local/share/uv/python
ENTRYPOINT ["entrypoints/extract_cpu_mineru_worker.sh"]


FROM runtime-base AS extract-gpu-mineru-build
FROM runtime-base AS extract-gpu-mineru-builder
# Install deps first to optimize layer cache
RUN --mount=type=cache,target=/root/.cache/uv \
--mount=type=bind,source=uv.dist.lock,target=uv.lock \
Expand All @@ -162,7 +164,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \
ADD uv.dist.lock ./uv.lock
ADD pyproject.toml README.md ./
ADD extract_worker ./extract_worker/
ADD entrypoints/extract_gpu_worker.sh ./entrypoints/extract_gpu_worker.sh
ADD entrypoints/extract_gpu_mineru_worker.sh ./entrypoints/extract_gpu_mineru_worker.sh

# Then install service
RUN --mount=type=cache,target=/root/.cache/uv uv sync -v --frozen --no-editable --extra mineru --extra gpu
Expand All @@ -171,5 +173,6 @@ RUN rm -rf ~/.cache/pip
# Slim runtime (as extract-cpu-worker): no toolchain, keep tesseract
FROM runtime-base AS extract-gpu-mineru-worker

COPY --from=extract-gpu-mineru-build /app /app
COPY --from=extract-gpu-mineru-builder /app /app
COPY --from=extract-gpu-mineru-builder /root/.local/share/uv/python /root/.local/share/uv/python
ENTRYPOINT ["entrypoints/extract_gpu_mineru_worker.sh"]
2 changes: 1 addition & 1 deletion workers/extract-worker/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ authors = [
readme = "README.md"
requires-python = ">=3.13,<3.15"
dependencies = [
"datashare-python~=0.9.0",
"datashare-python~=0.9.9",
"extract-core==0.7.1",
"temporalio==1.23.0",
]
Expand Down
26 changes: 23 additions & 3 deletions workers/translation-worker/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -34,12 +34,12 @@ RUN rm -rf ~/.cache/pip
ENTRYPOINT ["entrypoints/io_worker.sh"]


FROM translation-worker-builder AS inference-worker
FROM translation-worker-builder AS inference-c2translate-worker
# Install deps first to optimize layer cache
RUN --mount=type=cache,target=~/.cache/uv \
--mount=type=bind,source=uv.dist.lock,target=uv.lock \
--mount=type=bind,source=pyproject.toml,target=pyproject.toml \
uv sync -v --frozen --no-editable --no-install-project --extra inference
uv sync -v --frozen --no-editable --no-install-project --extra c2translate

# Then copy code
ADD uv.dist.lock ./uv.lock
Expand All @@ -48,7 +48,27 @@ ADD translation_worker ./translation_worker/
ADD entrypoints/inference_worker.sh ./entrypoints/inference_worker.sh

# Then install service
RUN --mount=type=cache,target=~/.cache/uv uv sync -v --frozen --no-editable --extra inference
RUN --mount=type=cache,target=~/.cache/uv uv sync -v --frozen --no-editable --extra c2translate
RUN rm -rf ~/.cache/pip

ENTRYPOINT ["entrypoints/inference_worker.sh"]


FROM translation-worker-builder AS inference-torch-worker
# Install deps first to optimize layer cache
RUN --mount=type=cache,target=~/.cache/uv \
--mount=type=bind,source=uv.dist.lock,target=uv.lock \
--mount=type=bind,source=pyproject.toml,target=pyproject.toml \
uv sync -v --frozen --no-editable --no-install-project --extra torch

# Then copy code
ADD uv.dist.lock ./uv.lock
ADD pyproject.toml README.md ./
ADD translation_worker ./translation_worker/
ADD entrypoints/inference_worker.sh ./entrypoints/inference_worker.sh

# Then install service
RUN --mount=type=cache,target=~/.cache/uv uv sync -v --frozen --no-editable --extra torch
RUN rm -rf ~/.cache/pip

ENTRYPOINT ["entrypoints/inference_worker.sh"]
12 changes: 6 additions & 6 deletions workers/translation-worker/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -29,14 +29,14 @@ dependencies = "translation_worker.dependencies:REGISTRY"
worker_config_cls = "translation_worker.config:WORKER_CONFIG_CLS"

[project.optional-dependencies]
inference = ['datashare-translation-worker[argos_inference]']
argos_inference = [
inference = ['datashare-translation-worker[c2translate]']
c2translate = [
"argostranslate==1.11.0",
]
hunyuan_inference = [
"torch>=2.11.0",
"transformers>=5.12.1",
"accelerate>=1.14.0",
torch = [
"torch==2.11.0",
"transformers==5.12.1",
"accelerate==1.14.0",
]

[tool.uv.sources]
Expand Down
2 changes: 1 addition & 1 deletion workers/translation-worker/tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,7 +194,7 @@ async def translation_inference_worker(
worker_id = f"test-translation-cpu-worker-{uuid.uuid4()}"
create_translation_batches = TranslationActivities(temporal_client=client)
translation_activities = [create_translation_batches.translate_docs]
task_queue = TaskQueue.INFERENCE
task_queue = TaskQueue.C2TRANSLATE_INFERENCE
worker_ctx = worker_context(
worker_id,
activities=translation_activities,
Expand Down
27 changes: 27 additions & 0 deletions workers/translation-worker/tests/test_constants.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
import pytest
from translation_worker.config import (
ArgosTranslatorConfig,
HunyuanMtTranslatorConfig,
TranslationConfig,
)
from translation_worker.constants import TaskQueue


@pytest.mark.parametrize(
("config", "expected_queue"),
[
(
TranslationConfig(translator=ArgosTranslatorConfig()),
TaskQueue.C2TRANSLATE_INFERENCE,
),
(
TranslationConfig(translator=HunyuanMtTranslatorConfig()),
TaskQueue.TORCH_INFERENCE,
),
],
)
def test_inference_queue(config: TranslationConfig, expected_queue: TaskQueue) -> None:
# When
inference_queue = TaskQueue.inference_queue(config)
# Then
assert inference_queue == expected_queue
21 changes: 11 additions & 10 deletions workers/translation-worker/translation_worker/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,14 +7,13 @@
from icij_common.registrable import RegistrableConfig
from pydantic import Discriminator, Field

from translation_worker.objects import (
from .objects import (
ArgosSentencizer,
SentenceSplitterModel,
TorchDevice,
TranslationModel,
)

from .constants import TorchDevice

if TYPE_CHECKING:
from translation_worker.processors import SentenceSplitter, Translator

Expand Down Expand Up @@ -72,7 +71,7 @@ class ArgosSentenceSplitterConfig(SentenceSplitterConfig):
sentencizer: ArgosSentencizer = ArgosSentencizer.MINI_SBD


class TranslatorConfig(_BaseProcessorConfig):
class BaseTranslatorConfig(_BaseProcessorConfig):
registry_key: ClassVar[str] = Field(frozen=True, default="model")
model: ClassVar[TranslationModel]

Expand All @@ -83,7 +82,7 @@ class TranslatorConfig(_BaseProcessorConfig):
splitter_discriminator = make_enum_discriminator("model", SentenceSplitterModel)


class ArgosTranslatorConfig(TranslatorConfig):
class ArgosTranslatorConfig(BaseTranslatorConfig):
model: ClassVar[TranslationModel] = Field(
frozen=True, default=TranslationModel.ARGOS
)
Expand All @@ -92,8 +91,10 @@ class ArgosTranslatorConfig(TranslatorConfig):
length_penalty: float = 0.2


class HunyuanMtTranslatorConfig(TranslatorConfig):
model_config = {"arbitrary_types_allowed": True}
DEFAULT_HUNYUAN_MODEL_REF = "tencent/Hunyuan-MT-Chimera-7B"


class HunyuanMtTranslatorConfig(BaseTranslatorConfig):
model: ClassVar[TranslationModel] = Field(
frozen=True, default=TranslationModel.HUNYUAN
)
Expand All @@ -108,8 +109,8 @@ class HunyuanMtTranslatorConfig(TranslatorConfig):
device_map: str = "auto"


_TranslatorConfig = tagged_union(
TranslatorConfig.__subclasses__(), lambda t: t.model.default
TranslatorConfig = tagged_union(
BaseTranslatorConfig.__subclasses__(), lambda t: t.model.default
)
translator_discriminator = make_enum_discriminator("model", TranslationModel)

Expand All @@ -119,7 +120,7 @@ class TranslationConfig(DatashareModel):
discriminator=Discriminator(splitter_discriminator),
default_factory=DefaultSentenceSplitterConfig,
)
translator: _TranslatorConfig = Field(
translator: TranslatorConfig = Field(
discriminator=Discriminator(translator_discriminator),
default_factory=ArgosTranslatorConfig,
)
Expand Down
23 changes: 17 additions & 6 deletions workers/translation-worker/translation_worker/constants.py
Original file line number Diff line number Diff line change
@@ -1,17 +1,28 @@
from enum import StrEnum
from typing import Self

from icij_common.es import DOC_CONTENT, DOC_LANGUAGE, DOC_ROOT_ID

from .config import TranslationConfig
from .objects import TranslationModel


class TaskQueue(StrEnum):
WORKFLOWS = "datashare.workflows"
IO = "translation.io"
INFERENCE = "translation.inference"


class TorchDevice(StrEnum):
CPU = "cpu"
GPU = "cuda"
TORCH_INFERENCE = "translation.inference.torch"
C2TRANSLATE_INFERENCE = "translation.inference.c2translate"

@classmethod
def inference_queue(cls, config: TranslationConfig) -> Self:
model = config.translator.model.default
match model:
case TranslationModel.ARGOS:
return TaskQueue.C2TRANSLATE_INFERENCE
case TranslationModel.HUNYUAN:
return TaskQueue.TORCH_INFERENCE
case _:
raise ValueError(f"unknown translation model {model}")


TRANSLATION_TASK_NAME = "translation"
Expand Down
5 changes: 5 additions & 0 deletions workers/translation-worker/translation_worker/objects.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,11 @@
from translation_worker.objects import TranslationConfig


class TorchDevice(StrEnum):
CPU = "cpu"
GPU = "cuda"


class SentenceSplitterModel(StrEnum):
ARGOS = "ARGOS"
DEFAULT = "DEFAULT"
Expand Down
4 changes: 2 additions & 2 deletions workers/translation-worker/translation_worker/processors.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from .config import TranslationWorkerConfig

if TYPE_CHECKING:
from .config import TranslatorConfig
from .config import BaseTranslatorConfig
from .objects import Language

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -43,7 +43,7 @@ def __exit__(self, exc_type, exc_val, exc_tb): ... # noqa: ANN001


class Translator(RegistrableFromConfig):
def __init__(self, config: "TranslatorConfig"):
def __init__(self, config: "BaseTranslatorConfig"):
self._config = config

self._source: Language | None = None
Expand Down
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
try:
from .argos import ArgosTranslator
except ImportError:
except ModuleNotFoundError:
ArgosTranslator = None

try:
from .hunyuan import HunyuanMtTranslator
except ImportError:
except ModuleNotFoundError:
HunyuanMtTranslator = None
2 changes: 1 addition & 1 deletion workers/translation-worker/translation_worker/utils.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from .constants import TorchDevice
from .objects import TorchDevice


def find_device(device_name: str = TorchDevice.CPU) -> TorchDevice.CPU:
Expand Down
Loading
Loading