claudevdm commented on code in PR #40239: URL: https://github.com/apache/beam/pull/40239#discussion_r4147540435
########## sdks/python/apache_beam/ml/inference/openai_inference.py: ########## @@ -0,0 +1,371 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""A ModelHandler for OpenAI models using the OpenAI Python SDK. + +This module provides an integration between Apache Beam's RunInference +transform and OpenAI's API, enabling batch inference and embeddings in +Beam pipelines. + +Example usage:: + + import apache_beam as beam + from apache_beam.ml.inference.base import RunInference + from apache_beam.ml.inference.openai_inference import ( + OpenAIModelHandler, + chat_completion_from_string, + ) + + # Basic text generation with chat completions + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + ) + + # With system prompt and structured output + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + system='You are a helpful assistant that responds concisely.', + response_format={ + 'type': 'json_schema', + 'json_schema': { + 'name': 'answer_response', + 'schema': { + 'type': 'object', + 'properties': { + 'answer': {'type': 'string'}, + 'confidence': {'type': 'number'}, + }, + 'required': ['answer', 'confidence'], + 'additionalProperties': False, + }, + 'strict': True, + }, + }, + ) + + with beam.Pipeline() as p: + results = ( + p + | beam.Create(['What is Apache Beam?', 'Explain MapReduce.']) + | RunInference(model_handler) + ) +""" + +import logging +from collections.abc import Callable +from collections.abc import Iterable +from collections.abc import Sequence +from typing import Any +from typing import Optional +from typing import Union Review Comment: Unused? ########## sdks/python/apache_beam/ml/inference/openai_inference.py: ########## @@ -0,0 +1,371 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""A ModelHandler for OpenAI models using the OpenAI Python SDK. + +This module provides an integration between Apache Beam's RunInference +transform and OpenAI's API, enabling batch inference and embeddings in +Beam pipelines. + +Example usage:: + + import apache_beam as beam + from apache_beam.ml.inference.base import RunInference + from apache_beam.ml.inference.openai_inference import ( + OpenAIModelHandler, + chat_completion_from_string, + ) + + # Basic text generation with chat completions + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + ) + + # With system prompt and structured output + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + system='You are a helpful assistant that responds concisely.', + response_format={ + 'type': 'json_schema', + 'json_schema': { + 'name': 'answer_response', + 'schema': { + 'type': 'object', + 'properties': { + 'answer': {'type': 'string'}, + 'confidence': {'type': 'number'}, + }, + 'required': ['answer', 'confidence'], + 'additionalProperties': False, + }, + 'strict': True, + }, + }, + ) + + with beam.Pipeline() as p: + results = ( + p + | beam.Create(['What is Apache Beam?', 'Explain MapReduce.']) + | RunInference(model_handler) + ) +""" + +import logging +from collections.abc import Callable +from collections.abc import Iterable +from collections.abc import Sequence +from typing import Any +from typing import Optional +from typing import Union + +from openai import APIConnectionError +from openai import APIStatusError +from openai import OpenAI + +from apache_beam.ml.inference import utils +from apache_beam.ml.inference.base import PredictionResult +from apache_beam.ml.inference.base import RemoteModelHandler + +__all__ = [ + 'OpenAIModelHandler', + 'chat_completion_from_string', + 'chat_completion_from_conversation', + 'embedding_from_string', +] + +LOGGER = logging.getLogger("OpenAIModelHandler") + + +def _retry_on_appropriate_error(exception: Exception) -> bool: + """Retry filter that returns True for retriable OpenAI API errors. + + Retries on HTTP 429 (rate limiting), HTTP 5xx (server errors), and + connection / timeout errors. + + Args: + exception: the exception encountered during the request/response loop. + + Returns: + True if the exception is retriable (429, 5xx, or connection error), + False otherwise. + """ + if isinstance(exception, APIConnectionError): + return True + if isinstance(exception, APIStatusError): + return exception.status_code == 429 or exception.status_code >= 500 Review Comment: Consider creating the client with max_retries=0 by default (unless the user sets it in client_args) so Beam's retry and throttling is the only layer, or at least document how the two interact. ########## sdks/python/apache_beam/ml/inference/openai_inference.py: ########## @@ -0,0 +1,371 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""A ModelHandler for OpenAI models using the OpenAI Python SDK. + +This module provides an integration between Apache Beam's RunInference +transform and OpenAI's API, enabling batch inference and embeddings in +Beam pipelines. + +Example usage:: + + import apache_beam as beam + from apache_beam.ml.inference.base import RunInference + from apache_beam.ml.inference.openai_inference import ( + OpenAIModelHandler, + chat_completion_from_string, + ) + + # Basic text generation with chat completions + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + ) + + # With system prompt and structured output + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + system='You are a helpful assistant that responds concisely.', + response_format={ + 'type': 'json_schema', + 'json_schema': { + 'name': 'answer_response', + 'schema': { + 'type': 'object', + 'properties': { + 'answer': {'type': 'string'}, + 'confidence': {'type': 'number'}, + }, + 'required': ['answer', 'confidence'], + 'additionalProperties': False, + }, + 'strict': True, + }, + }, + ) + + with beam.Pipeline() as p: + results = ( + p + | beam.Create(['What is Apache Beam?', 'Explain MapReduce.']) + | RunInference(model_handler) + ) +""" + +import logging +from collections.abc import Callable +from collections.abc import Iterable +from collections.abc import Sequence +from typing import Any +from typing import Optional +from typing import Union + +from openai import APIConnectionError +from openai import APIStatusError +from openai import OpenAI + +from apache_beam.ml.inference import utils +from apache_beam.ml.inference.base import PredictionResult +from apache_beam.ml.inference.base import RemoteModelHandler + +__all__ = [ + 'OpenAIModelHandler', + 'chat_completion_from_string', + 'chat_completion_from_conversation', + 'embedding_from_string', +] + +LOGGER = logging.getLogger("OpenAIModelHandler") + + +def _retry_on_appropriate_error(exception: Exception) -> bool: + """Retry filter that returns True for retriable OpenAI API errors. + + Retries on HTTP 429 (rate limiting), HTTP 5xx (server errors), and + connection / timeout errors. + + Args: + exception: the exception encountered during the request/response loop. + + Returns: + True if the exception is retriable (429, 5xx, or connection error), + False otherwise. + """ + if isinstance(exception, APIConnectionError): + return True + if isinstance(exception, APIStatusError): + return exception.status_code == 429 or exception.status_code >= 500 + return False + + +def chat_completion_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string prompts to OpenAI's Chat Completions API. + + Each string in the batch is sent as a user message. If a 'system' parameter + is provided in inference_args, a system message is prepended. The results + are returned as a list of ChatCompletion response objects. + + Args: + model_name: the OpenAI model to use (e.g. 'gpt-4o', 'gpt-4o-mini'). + batch: the string prompts to send to OpenAI. + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call (e.g. 'temperature', 'max_tokens', 'response_format', 'system'). + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for prompt in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.append({"role": "user", "content": prompt}) + response = client.chat.completions.create( Review Comment: This sends prompts one at a time. Consider sending prompts in batch concurrently like https://github.com/apache/beam/blob/b4623815617ef1fc3a2fbfd0bd4d5f96c2402a88/sdks/python/apache_beam/ml/inference/vllm_inference.py#L460-L465 ########## sdks/python/apache_beam/ml/inference/openai_inference.py: ########## @@ -0,0 +1,371 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""A ModelHandler for OpenAI models using the OpenAI Python SDK. + +This module provides an integration between Apache Beam's RunInference +transform and OpenAI's API, enabling batch inference and embeddings in +Beam pipelines. + +Example usage:: + + import apache_beam as beam + from apache_beam.ml.inference.base import RunInference + from apache_beam.ml.inference.openai_inference import ( + OpenAIModelHandler, + chat_completion_from_string, + ) + + # Basic text generation with chat completions + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + ) + + # With system prompt and structured output + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + system='You are a helpful assistant that responds concisely.', + response_format={ + 'type': 'json_schema', + 'json_schema': { + 'name': 'answer_response', + 'schema': { + 'type': 'object', + 'properties': { + 'answer': {'type': 'string'}, + 'confidence': {'type': 'number'}, + }, + 'required': ['answer', 'confidence'], + 'additionalProperties': False, + }, + 'strict': True, + }, + }, + ) + + with beam.Pipeline() as p: + results = ( + p + | beam.Create(['What is Apache Beam?', 'Explain MapReduce.']) + | RunInference(model_handler) + ) +""" + +import logging +from collections.abc import Callable +from collections.abc import Iterable +from collections.abc import Sequence +from typing import Any +from typing import Optional +from typing import Union + +from openai import APIConnectionError +from openai import APIStatusError +from openai import OpenAI + +from apache_beam.ml.inference import utils +from apache_beam.ml.inference.base import PredictionResult +from apache_beam.ml.inference.base import RemoteModelHandler + +__all__ = [ + 'OpenAIModelHandler', + 'chat_completion_from_string', + 'chat_completion_from_conversation', + 'embedding_from_string', +] + +LOGGER = logging.getLogger("OpenAIModelHandler") + + +def _retry_on_appropriate_error(exception: Exception) -> bool: + """Retry filter that returns True for retriable OpenAI API errors. + + Retries on HTTP 429 (rate limiting), HTTP 5xx (server errors), and + connection / timeout errors. + + Args: + exception: the exception encountered during the request/response loop. + + Returns: + True if the exception is retriable (429, 5xx, or connection error), + False otherwise. + """ + if isinstance(exception, APIConnectionError): + return True + if isinstance(exception, APIStatusError): + return exception.status_code == 429 or exception.status_code >= 500 + return False + + +def chat_completion_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string prompts to OpenAI's Chat Completions API. + + Each string in the batch is sent as a user message. If a 'system' parameter + is provided in inference_args, a system message is prepended. The results + are returned as a list of ChatCompletion response objects. + + Args: + model_name: the OpenAI model to use (e.g. 'gpt-4o', 'gpt-4o-mini'). + batch: the string prompts to send to OpenAI. + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call (e.g. 'temperature', 'max_tokens', 'response_format', 'system'). + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for prompt in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.append({"role": "user", "content": prompt}) + response = client.chat.completions.create( + model=model_name, messages=messages, **inf_args) + responses.append(response) + return responses + + +def chat_completion_from_conversation( + model_name: str, + batch: Sequence[list[dict[str, Any]]], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends multi-turn conversations to OpenAI. + + Each element in the batch is a list of message dicts (e.g. with 'role' and + 'content' keys), representing a multi-turn conversation. If a 'system' + parameter is provided in inference_args, a system message is prepended. + + Args: + model_name: the OpenAI model to use. + batch: a sequence of conversations (each a list of message dicts). + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call. + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for conversation in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.extend(conversation) + response = client.chat.completions.create( + model=model_name, messages=messages, **inf_args) + responses.append(response) + return responses + + +def embedding_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string inputs to OpenAI's Embeddings API. + + The batch of string inputs is sent to client.embeddings.create. The returned + embeddings are sorted by their index to guarantee ordering matches the batch. + + Args: + model_name: the OpenAI embedding model to use (e.g. 'text-embedding-3-small'). + batch: the string inputs to embed. + client: the OpenAI client instance. + inference_args: additional arguments passed to the embeddings.create call + (e.g. 'dimensions', 'encoding_format'). + + Returns: + A list of Embedding objects matching the batch order. + """ + inf_args = dict(inference_args) + response = client.embeddings.create( + model=model_name, input=list(batch), **inf_args) + sorted_data = sorted(response.data, key=lambda x: x.index) + return sorted_data + + +class OpenAIModelHandler(RemoteModelHandler[Any, PredictionResult, OpenAI]): + def __init__( + self, + model_name: str, + request_fn: Callable[[str, Sequence[Any], OpenAI, dict[str, Any]], Any], + api_key: Optional[str] = None, + *, + organization: Optional[str] = None, + project: Optional[str] = None, + base_url: Optional[str] = None, + client_args: Optional[dict[str, Any]] = None, + system: Optional[str] = None, + response_format: Optional[dict[str, Any]] = None, + min_batch_size: Optional[int] = None, + max_batch_size: Optional[int] = None, + max_batch_duration_secs: Optional[int] = None, + max_batch_weight: Optional[int] = None, + element_size_fn: Optional[Callable[[Any], int]] = None, + batch_length_fn: Optional[Callable[[Any], int]] = None, + batch_bucket_boundaries: Optional[list[int]] = None, + **kwargs): + """Implementation of the ModelHandler interface for OpenAI models. + + **NOTE:** This API and its implementation are under development and + do not provide backward compatibility guarantees. + + This handler connects to the OpenAI API using the OpenAI Python SDK + to run inference using models such as GPT-4o, GPT-4o-mini, or embedding + models. It supports chat completions from string prompts or multi-turn + conversations, embeddings, system prompts, structured outputs, and + custom OpenAI-compatible endpoints via `base_url`. + + Args: + model_name: the OpenAI model to send requests to (e.g. + 'gpt-4o', 'gpt-4o-mini', 'text-embedding-3-small'). + request_fn: the function to use to send requests. Should take the Review Comment: Could we spell out the contract here: request_fn must return a list with exactly one response per input, in the same order? ########## sdks/python/apache_beam/ml/inference/openai_inference.py: ########## @@ -0,0 +1,371 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""A ModelHandler for OpenAI models using the OpenAI Python SDK. + +This module provides an integration between Apache Beam's RunInference +transform and OpenAI's API, enabling batch inference and embeddings in +Beam pipelines. + +Example usage:: + + import apache_beam as beam + from apache_beam.ml.inference.base import RunInference + from apache_beam.ml.inference.openai_inference import ( + OpenAIModelHandler, + chat_completion_from_string, + ) + + # Basic text generation with chat completions + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + ) + + # With system prompt and structured output + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + system='You are a helpful assistant that responds concisely.', + response_format={ + 'type': 'json_schema', + 'json_schema': { + 'name': 'answer_response', + 'schema': { + 'type': 'object', + 'properties': { + 'answer': {'type': 'string'}, + 'confidence': {'type': 'number'}, + }, + 'required': ['answer', 'confidence'], + 'additionalProperties': False, + }, + 'strict': True, + }, + }, + ) + + with beam.Pipeline() as p: + results = ( + p + | beam.Create(['What is Apache Beam?', 'Explain MapReduce.']) + | RunInference(model_handler) + ) +""" + +import logging +from collections.abc import Callable +from collections.abc import Iterable +from collections.abc import Sequence +from typing import Any +from typing import Optional +from typing import Union + +from openai import APIConnectionError +from openai import APIStatusError +from openai import OpenAI + +from apache_beam.ml.inference import utils +from apache_beam.ml.inference.base import PredictionResult +from apache_beam.ml.inference.base import RemoteModelHandler + +__all__ = [ + 'OpenAIModelHandler', + 'chat_completion_from_string', + 'chat_completion_from_conversation', + 'embedding_from_string', +] + +LOGGER = logging.getLogger("OpenAIModelHandler") + + +def _retry_on_appropriate_error(exception: Exception) -> bool: + """Retry filter that returns True for retriable OpenAI API errors. + + Retries on HTTP 429 (rate limiting), HTTP 5xx (server errors), and + connection / timeout errors. + + Args: + exception: the exception encountered during the request/response loop. + + Returns: + True if the exception is retriable (429, 5xx, or connection error), + False otherwise. + """ + if isinstance(exception, APIConnectionError): + return True + if isinstance(exception, APIStatusError): + return exception.status_code == 429 or exception.status_code >= 500 + return False + + +def chat_completion_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string prompts to OpenAI's Chat Completions API. + + Each string in the batch is sent as a user message. If a 'system' parameter + is provided in inference_args, a system message is prepended. The results + are returned as a list of ChatCompletion response objects. + + Args: + model_name: the OpenAI model to use (e.g. 'gpt-4o', 'gpt-4o-mini'). + batch: the string prompts to send to OpenAI. + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call (e.g. 'temperature', 'max_tokens', 'response_format', 'system'). + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for prompt in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.append({"role": "user", "content": prompt}) + response = client.chat.completions.create( + model=model_name, messages=messages, **inf_args) + responses.append(response) + return responses + + +def chat_completion_from_conversation( + model_name: str, + batch: Sequence[list[dict[str, Any]]], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends multi-turn conversations to OpenAI. + + Each element in the batch is a list of message dicts (e.g. with 'role' and + 'content' keys), representing a multi-turn conversation. If a 'system' + parameter is provided in inference_args, a system message is prepended. + + Args: + model_name: the OpenAI model to use. + batch: a sequence of conversations (each a list of message dicts). + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call. + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for conversation in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.extend(conversation) + response = client.chat.completions.create( + model=model_name, messages=messages, **inf_args) + responses.append(response) + return responses + + +def embedding_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string inputs to OpenAI's Embeddings API. + + The batch of string inputs is sent to client.embeddings.create. The returned + embeddings are sorted by their index to guarantee ordering matches the batch. + + Args: + model_name: the OpenAI embedding model to use (e.g. 'text-embedding-3-small'). + batch: the string inputs to embed. + client: the OpenAI client instance. + inference_args: additional arguments passed to the embeddings.create call + (e.g. 'dimensions', 'encoding_format'). + + Returns: + A list of Embedding objects matching the batch order. + """ + inf_args = dict(inference_args) + response = client.embeddings.create( Review Comment: Should we set a default max batch size so that we dont exceed the api batch size limits? ########## sdks/python/apache_beam/ml/inference/openai_inference.py: ########## @@ -0,0 +1,371 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +"""A ModelHandler for OpenAI models using the OpenAI Python SDK. + +This module provides an integration between Apache Beam's RunInference +transform and OpenAI's API, enabling batch inference and embeddings in +Beam pipelines. + +Example usage:: + + import apache_beam as beam + from apache_beam.ml.inference.base import RunInference + from apache_beam.ml.inference.openai_inference import ( + OpenAIModelHandler, + chat_completion_from_string, + ) + + # Basic text generation with chat completions + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + ) + + # With system prompt and structured output + model_handler = OpenAIModelHandler( + model_name='gpt-4o-mini', + api_key='your-api-key', + request_fn=chat_completion_from_string, + system='You are a helpful assistant that responds concisely.', + response_format={ + 'type': 'json_schema', + 'json_schema': { + 'name': 'answer_response', + 'schema': { + 'type': 'object', + 'properties': { + 'answer': {'type': 'string'}, + 'confidence': {'type': 'number'}, + }, + 'required': ['answer', 'confidence'], + 'additionalProperties': False, + }, + 'strict': True, + }, + }, + ) + + with beam.Pipeline() as p: + results = ( + p + | beam.Create(['What is Apache Beam?', 'Explain MapReduce.']) + | RunInference(model_handler) + ) +""" + +import logging +from collections.abc import Callable +from collections.abc import Iterable +from collections.abc import Sequence +from typing import Any +from typing import Optional +from typing import Union + +from openai import APIConnectionError +from openai import APIStatusError +from openai import OpenAI + +from apache_beam.ml.inference import utils +from apache_beam.ml.inference.base import PredictionResult +from apache_beam.ml.inference.base import RemoteModelHandler + +__all__ = [ + 'OpenAIModelHandler', + 'chat_completion_from_string', + 'chat_completion_from_conversation', + 'embedding_from_string', +] + +LOGGER = logging.getLogger("OpenAIModelHandler") + + +def _retry_on_appropriate_error(exception: Exception) -> bool: + """Retry filter that returns True for retriable OpenAI API errors. + + Retries on HTTP 429 (rate limiting), HTTP 5xx (server errors), and + connection / timeout errors. + + Args: + exception: the exception encountered during the request/response loop. + + Returns: + True if the exception is retriable (429, 5xx, or connection error), + False otherwise. + """ + if isinstance(exception, APIConnectionError): + return True + if isinstance(exception, APIStatusError): + return exception.status_code == 429 or exception.status_code >= 500 + return False + + +def chat_completion_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string prompts to OpenAI's Chat Completions API. + + Each string in the batch is sent as a user message. If a 'system' parameter + is provided in inference_args, a system message is prepended. The results + are returned as a list of ChatCompletion response objects. + + Args: + model_name: the OpenAI model to use (e.g. 'gpt-4o', 'gpt-4o-mini'). + batch: the string prompts to send to OpenAI. + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call (e.g. 'temperature', 'max_tokens', 'response_format', 'system'). + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for prompt in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.append({"role": "user", "content": prompt}) + response = client.chat.completions.create( + model=model_name, messages=messages, **inf_args) + responses.append(response) + return responses + + +def chat_completion_from_conversation( + model_name: str, + batch: Sequence[list[dict[str, Any]]], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends multi-turn conversations to OpenAI. + + Each element in the batch is a list of message dicts (e.g. with 'role' and + 'content' keys), representing a multi-turn conversation. If a 'system' + parameter is provided in inference_args, a system message is prepended. + + Args: + model_name: the OpenAI model to use. + batch: a sequence of conversations (each a list of message dicts). + client: the OpenAI client instance. + inference_args: additional arguments passed to the chat.completions.create + call. + """ + inf_args = dict(inference_args) + system = inf_args.pop('system', None) + responses = [] + for conversation in batch: + messages: list[dict[str, Any]] = [] + if system is not None: + messages.append({"role": "system", "content": system}) + messages.extend(conversation) + response = client.chat.completions.create( + model=model_name, messages=messages, **inf_args) + responses.append(response) + return responses + + +def embedding_from_string( + model_name: str, + batch: Sequence[str], + client: OpenAI, + inference_args: dict[str, Any]) -> list[Any]: + """Request function that sends string inputs to OpenAI's Embeddings API. + + The batch of string inputs is sent to client.embeddings.create. The returned + embeddings are sorted by their index to guarantee ordering matches the batch. + + Args: + model_name: the OpenAI embedding model to use (e.g. 'text-embedding-3-small'). + batch: the string inputs to embed. + client: the OpenAI client instance. + inference_args: additional arguments passed to the embeddings.create call + (e.g. 'dimensions', 'encoding_format'). + + Returns: + A list of Embedding objects matching the batch order. + """ + inf_args = dict(inference_args) + response = client.embeddings.create( + model=model_name, input=list(batch), **inf_args) + sorted_data = sorted(response.data, key=lambda x: x.index) + return sorted_data + + +class OpenAIModelHandler(RemoteModelHandler[Any, PredictionResult, OpenAI]): + def __init__( + self, + model_name: str, + request_fn: Callable[[str, Sequence[Any], OpenAI, dict[str, Any]], Any], + api_key: Optional[str] = None, + *, + organization: Optional[str] = None, + project: Optional[str] = None, + base_url: Optional[str] = None, + client_args: Optional[dict[str, Any]] = None, + system: Optional[str] = None, Review Comment: Should we take this off the handler? Just let users set it in inference args instead? -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
