From cff7b1738e9857137148c13d74e4891b4b22936a Mon Sep 17 00:00:00 2001 From: Tedi Papajorgji Date: Thu, 3 Mar 2022 14:14:36 -0500 Subject: [PATCH 1/4] Expose device_requests to DockerOperator --- airflow/providers/docker/operators/docker.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/airflow/providers/docker/operators/docker.py b/airflow/providers/docker/operators/docker.py index 4eb2560fc54b3..126636936818f 100644 --- a/airflow/providers/docker/operators/docker.py +++ b/airflow/providers/docker/operators/docker.py @@ -128,6 +128,8 @@ class DockerOperator(BaseOperator): file before manually shutting down the image. Useful for cases where users want a pickle serialized output that is not posted to logs :param retrieve_output_path: path for output file that will be retrieved and passed to xcom + :param device_requests: Expose host resources such as GPUs to the container, as a list of docker.types.DeviceRequest instances. + To expose all GPU's to the container you may use device_requests=[docker.types.DeviceRequest(count=-1, capabilities=[['gpu']])] """ template_fields: Sequence[str] = ('image', 'command', 'environment', 'container_name') @@ -174,6 +176,7 @@ def __init__( extra_hosts: Optional[Dict[str, str]] = None, retrieve_output: bool = False, retrieve_output_path: Optional[str] = None, + device_requests: Optional(List[docker.types.DeviceRequest]) = None, **kwargs, ) -> None: super().__init__(**kwargs) @@ -217,6 +220,7 @@ def __init__( self.container = None self.retrieve_output = retrieve_output self.retrieve_output_path = retrieve_output_path + self.device_requests = device_requests def get_hook(self) -> DockerHook: """ @@ -279,6 +283,7 @@ def _run_image_with_mounts( cap_add=self.cap_add, extra_hosts=self.extra_hosts, privileged=self.privileged, + device_requests=self.device_requests, ), image=self.image, user=self.user, From 32571018767868c9aba848ca5b7f5f34594b7855 Mon Sep 17 00:00:00 2001 From: Tedi Papajorgji Date: Thu, 3 Mar 2022 14:25:05 -0500 Subject: [PATCH 2/4] Update some versioning stuff and changelog --- airflow/providers/docker/CHANGELOG.rst | 9 +++++++++ airflow/providers/docker/provider.yaml | 1 + 2 files changed, 10 insertions(+) diff --git a/airflow/providers/docker/CHANGELOG.rst b/airflow/providers/docker/CHANGELOG.rst index 73c2dd8ac7aa3..276434ef54949 100644 --- a/airflow/providers/docker/CHANGELOG.rst +++ b/airflow/providers/docker/CHANGELOG.rst @@ -19,6 +19,15 @@ Changelog --------- +2.4.2 +..... + +Features +~~~~~~~~ + +* ``Add support for device_requests. Used to expose host resources such as GPUs to the container (#21974)`` + + 2.4.1 ..... diff --git a/airflow/providers/docker/provider.yaml b/airflow/providers/docker/provider.yaml index fbf5533e49b4a..f8613c896e582 100644 --- a/airflow/providers/docker/provider.yaml +++ b/airflow/providers/docker/provider.yaml @@ -22,6 +22,7 @@ description: | `Docker `__ versions: + - 2.4.2 - 2.4.1 - 2.4.0 - 2.3.0 From a569f75a062c6d66172afc76b5995cc09835a110 Mon Sep 17 00:00:00 2001 From: Tedi Papajorgji Date: Thu, 3 Mar 2022 14:51:43 -0500 Subject: [PATCH 3/4] fix tests --- airflow/providers/docker/operators/docker.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/airflow/providers/docker/operators/docker.py b/airflow/providers/docker/operators/docker.py index 126636936818f..f92052179d334 100644 --- a/airflow/providers/docker/operators/docker.py +++ b/airflow/providers/docker/operators/docker.py @@ -25,7 +25,7 @@ from docker import APIClient, tls from docker.errors import APIError -from docker.types import Mount +from docker.types import Mount, DeviceRequest from airflow.exceptions import AirflowException from airflow.models import BaseOperator @@ -176,7 +176,7 @@ def __init__( extra_hosts: Optional[Dict[str, str]] = None, retrieve_output: bool = False, retrieve_output_path: Optional[str] = None, - device_requests: Optional(List[docker.types.DeviceRequest]) = None, + device_requests: Optional(List[DeviceRequest]) = None, **kwargs, ) -> None: super().__init__(**kwargs) From b3ae19137b63d1d86558d947bd618703fe9aea0c Mon Sep 17 00:00:00 2001 From: Tedi Papajorgji Date: Mon, 14 Mar 2022 18:11:06 -0400 Subject: [PATCH 4/4] fix line length & test failures --- airflow/providers/docker/operators/docker.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/airflow/providers/docker/operators/docker.py b/airflow/providers/docker/operators/docker.py index f92052179d334..34f34cfcd4f2e 100644 --- a/airflow/providers/docker/operators/docker.py +++ b/airflow/providers/docker/operators/docker.py @@ -25,7 +25,7 @@ from docker import APIClient, tls from docker.errors import APIError -from docker.types import Mount, DeviceRequest +from docker.types import DeviceRequest, Mount from airflow.exceptions import AirflowException from airflow.models import BaseOperator @@ -128,8 +128,9 @@ class DockerOperator(BaseOperator): file before manually shutting down the image. Useful for cases where users want a pickle serialized output that is not posted to logs :param retrieve_output_path: path for output file that will be retrieved and passed to xcom - :param device_requests: Expose host resources such as GPUs to the container, as a list of docker.types.DeviceRequest instances. - To expose all GPU's to the container you may use device_requests=[docker.types.DeviceRequest(count=-1, capabilities=[['gpu']])] + :param device_requests: Expose host resources such as GPUs to the container, as a list of docker.types.DeviceRequest + instances. To expose all GPU's to the container you may use + device_requests=[docker.types.DeviceRequest(count=-1, capabilities=[['gpu']])] """ template_fields: Sequence[str] = ('image', 'command', 'environment', 'container_name') @@ -176,7 +177,7 @@ def __init__( extra_hosts: Optional[Dict[str, str]] = None, retrieve_output: bool = False, retrieve_output_path: Optional[str] = None, - device_requests: Optional(List[DeviceRequest]) = None, + device_requests: Optional[List[DeviceRequest]] = None, **kwargs, ) -> None: super().__init__(**kwargs)