diff --git a/airflow/providers/docker/CHANGELOG.rst b/airflow/providers/docker/CHANGELOG.rst index e5adf853ee5af..ed7fbe27ccaa7 100644 --- a/airflow/providers/docker/CHANGELOG.rst +++ b/airflow/providers/docker/CHANGELOG.rst @@ -19,6 +19,14 @@ Changelog --------- +2.5.2 +..... + +Features +~~~~~~~~ + +* ``Add support for device_requests. Used to expose host resources such as GPUs to the container (#21974)`` + 2.5.1 ..... diff --git a/airflow/providers/docker/operators/docker.py b/airflow/providers/docker/operators/docker.py index 53dc5d92dea8d..e533728d2148b 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 DeviceRequest, Mount from airflow.exceptions import AirflowException from airflow.models import BaseOperator @@ -135,6 +135,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']])] """ template_fields: Sequence[str] = ('image', 'command', 'environment', 'container_name') @@ -181,6 +184,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, **kwargs, ) -> None: super().__init__(**kwargs) @@ -224,6 +228,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: """ @@ -286,6 +291,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, diff --git a/airflow/providers/docker/provider.yaml b/airflow/providers/docker/provider.yaml index 895961a5f14a3..da807840070ad 100644 --- a/airflow/providers/docker/provider.yaml +++ b/airflow/providers/docker/provider.yaml @@ -22,6 +22,7 @@ description: | `Docker `__ versions: + - 2.5.2 - 2.5.1 - 2.5.0 - 2.4.1