predictionguard 2.10.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- predictionguard/__init__.py +6 -0
- predictionguard/client.py +151 -0
- predictionguard/src/audio.py +194 -0
- predictionguard/src/chat.py +381 -0
- predictionguard/src/completions.py +244 -0
- predictionguard/src/detokenize.py +124 -0
- predictionguard/src/documents.py +145 -0
- predictionguard/src/embeddings.py +214 -0
- predictionguard/src/factuality.py +96 -0
- predictionguard/src/injection.py +100 -0
- predictionguard/src/models.py +101 -0
- predictionguard/src/pii.py +106 -0
- predictionguard/src/rerank.py +131 -0
- predictionguard/src/tokenize.py +119 -0
- predictionguard/src/toxicity.py +92 -0
- predictionguard/src/translate.py +25 -0
- predictionguard/version.py +2 -0
- predictionguard-2.10.0.dist-info/METADATA +51 -0
- predictionguard-2.10.0.dist-info/RECORD +21 -0
- predictionguard-2.10.0.dist-info/WHEEL +4 -0
- predictionguard-2.10.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,151 @@
|
|
|
1
|
+
import os
|
|
2
|
+
|
|
3
|
+
import requests
|
|
4
|
+
from typing import Optional, Union
|
|
5
|
+
|
|
6
|
+
from .src.audio import Audio
|
|
7
|
+
from .src.chat import Chat
|
|
8
|
+
from .src.completions import Completions
|
|
9
|
+
from .src.detokenize import Detokenize
|
|
10
|
+
from .src.documents import Documents
|
|
11
|
+
from .src.embeddings import Embeddings
|
|
12
|
+
from .src.rerank import Rerank
|
|
13
|
+
from .src.tokenize import Tokenize
|
|
14
|
+
from .src.translate import Translate
|
|
15
|
+
from .src.factuality import Factuality
|
|
16
|
+
from .src.toxicity import Toxicity
|
|
17
|
+
from .src.pii import Pii
|
|
18
|
+
from .src.injection import Injection
|
|
19
|
+
from .src.models import Models
|
|
20
|
+
from .version import __version__
|
|
21
|
+
|
|
22
|
+
__all__ = [
|
|
23
|
+
"PredictionGuard", "Chat", "Completions", "Embeddings",
|
|
24
|
+
"Audio", "Documents", "Rerank", "Tokenize", "Translate",
|
|
25
|
+
"Detokenize", "Factuality", "Toxicity", "Pii", "Injection",
|
|
26
|
+
"Models"
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
class PredictionGuard:
|
|
30
|
+
"""PredictionGuard provides access the Prediction Guard API."""
|
|
31
|
+
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
api_key: Optional[str] = None,
|
|
35
|
+
url: Optional[str] = None,
|
|
36
|
+
timeout: Optional[Union[int, float]] = None
|
|
37
|
+
) -> None:
|
|
38
|
+
"""
|
|
39
|
+
:param api_key: api_key represents PG api key.
|
|
40
|
+
:param url: url represents the transport and domain:port
|
|
41
|
+
:param timeout: request timeout in seconds.
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
# Get the access api_key.
|
|
45
|
+
if not api_key:
|
|
46
|
+
api_key = os.environ.get("PREDICTIONGUARD_API_KEY")
|
|
47
|
+
|
|
48
|
+
if not api_key:
|
|
49
|
+
raise ValueError(
|
|
50
|
+
"No api_key provided or in environment. "
|
|
51
|
+
"Please provide the api_key as "
|
|
52
|
+
"client = PredictionGuard(api_key=<your_api_key>) "
|
|
53
|
+
"or as PREDICTIONGUARD_API_KEY in your environment."
|
|
54
|
+
)
|
|
55
|
+
self.api_key = api_key
|
|
56
|
+
|
|
57
|
+
if not url:
|
|
58
|
+
url = os.environ.get("PREDICTIONGUARD_URL")
|
|
59
|
+
if not url:
|
|
60
|
+
url = "https://api.predictionguard.com"
|
|
61
|
+
self.url = url
|
|
62
|
+
|
|
63
|
+
if not timeout:
|
|
64
|
+
timeout = os.environ.get("TIMEOUT")
|
|
65
|
+
if not timeout:
|
|
66
|
+
timeout = None
|
|
67
|
+
if timeout:
|
|
68
|
+
try:
|
|
69
|
+
timeout = float(timeout)
|
|
70
|
+
except ValueError:
|
|
71
|
+
raise ValueError(
|
|
72
|
+
"Timeout must be of type integer or float, not %s." % (type(timeout).__name__,)
|
|
73
|
+
)
|
|
74
|
+
except TypeError:
|
|
75
|
+
raise TypeError(
|
|
76
|
+
"Timeout should be of type integer or float, not %s." % (type(timeout).__name__,)
|
|
77
|
+
)
|
|
78
|
+
self.timeout = timeout
|
|
79
|
+
|
|
80
|
+
# Connect to Prediction Guard and set the access api_key.
|
|
81
|
+
self._connect_client()
|
|
82
|
+
|
|
83
|
+
# Pass Prediction Guard class variables to inner classes
|
|
84
|
+
self.chat: Chat = Chat(self.api_key, self.url, self.timeout)
|
|
85
|
+
"""Chat generates chat completions based on a conversation history"""
|
|
86
|
+
|
|
87
|
+
self.completions: Completions = Completions(self.api_key, self.url, self.timeout)
|
|
88
|
+
"""Completions generates text completions based on the provided input"""
|
|
89
|
+
|
|
90
|
+
self.embeddings: Embeddings = Embeddings(self.api_key, self.url, self.timeout)
|
|
91
|
+
"""Embedding generates chat completions based on a conversation history."""
|
|
92
|
+
|
|
93
|
+
self.audio: Audio = Audio(self.api_key, self.url, self.timeout)
|
|
94
|
+
"""Audio allows for the transcription of audio files."""
|
|
95
|
+
|
|
96
|
+
self.documents: Documents = Documents(self.api_key, self.url, self.timeout)
|
|
97
|
+
"""Documents allows you to extract text from various document file types."""
|
|
98
|
+
|
|
99
|
+
self.rerank: Rerank = Rerank(self.api_key, self.url, self.timeout)
|
|
100
|
+
"""Rerank sorts text inputs by semantic relevance to a specified query."""
|
|
101
|
+
|
|
102
|
+
self.translate: Translate = Translate(self.api_key, self.url, self.timeout)
|
|
103
|
+
"""Translate converts text from one language to another."""
|
|
104
|
+
|
|
105
|
+
self.factuality: Factuality = Factuality(self.api_key, self.url, self.timeout)
|
|
106
|
+
"""Factuality checks the factuality of a given text compared to a reference."""
|
|
107
|
+
|
|
108
|
+
self.toxicity: Toxicity = Toxicity(self.api_key, self.url, self.timeout)
|
|
109
|
+
"""Toxicity checks the toxicity of a given text."""
|
|
110
|
+
|
|
111
|
+
self.pii: Pii = Pii(self.api_key, self.url, self.timeout)
|
|
112
|
+
"""Pii replaces personal information such as names, SSNs, and emails in a given text."""
|
|
113
|
+
|
|
114
|
+
self.injection: Injection = Injection(self.api_key, self.url, self.timeout)
|
|
115
|
+
"""Injection detects potential prompt injection attacks in a given prompt."""
|
|
116
|
+
|
|
117
|
+
self.tokenize: Tokenize = Tokenize(self.api_key, self.url, self.timeout)
|
|
118
|
+
"""Tokenize generates tokens for input text."""
|
|
119
|
+
|
|
120
|
+
self.detokenize: Detokenize = Detokenize(self.api_key, self.url, self.timeout)
|
|
121
|
+
"""Detokenizes generates text for input tokens."""
|
|
122
|
+
|
|
123
|
+
self.models: Models = Models(self.api_key, self.url, self.timeout)
|
|
124
|
+
"""Models lists all of the models available in the Prediction Guard API."""
|
|
125
|
+
|
|
126
|
+
def _connect_client(self) -> None:
|
|
127
|
+
|
|
128
|
+
# Prepare the proper headers.
|
|
129
|
+
headers = {
|
|
130
|
+
"Content-Type": "application/json",
|
|
131
|
+
"Authorization": "Bearer " + self.api_key,
|
|
132
|
+
"User-Agent": "Prediction Guard Python Client: " + __version__,
|
|
133
|
+
}
|
|
134
|
+
|
|
135
|
+
# Try listing models to make sure we can connect.
|
|
136
|
+
response = requests.request("GET", self.url + "/completions", headers=headers, timeout=self.timeout)
|
|
137
|
+
|
|
138
|
+
# If the connection was unsuccessful, raise an exception.
|
|
139
|
+
if response.status_code == 200:
|
|
140
|
+
pass
|
|
141
|
+
elif response.status_code == 401:
|
|
142
|
+
raise ValueError(
|
|
143
|
+
"Could not connect to Prediction Guard API with the given api_key. "
|
|
144
|
+
"Please check your access api_key and try again."
|
|
145
|
+
)
|
|
146
|
+
elif response.status_code == 404:
|
|
147
|
+
raise ValueError(
|
|
148
|
+
"Could not connect to Prediction Guard API with given url. "
|
|
149
|
+
"Please check url specified, if no url specified, "
|
|
150
|
+
"Please contact support."
|
|
151
|
+
)
|
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
from typing import Any, Dict, List, Optional
|
|
2
|
+
|
|
3
|
+
import requests
|
|
4
|
+
|
|
5
|
+
from ..version import __version__
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class Audio:
|
|
9
|
+
"""
|
|
10
|
+
Audio generates a response based on audio data.
|
|
11
|
+
|
|
12
|
+
Usage::
|
|
13
|
+
|
|
14
|
+
import os
|
|
15
|
+
import json
|
|
16
|
+
|
|
17
|
+
from predictionguard import PredictionGuard
|
|
18
|
+
|
|
19
|
+
# Set your Prediction Guard token and url as an environmental variable.
|
|
20
|
+
os.environ["PREDICTIONGUARD_API_KEY"] = "<api key>"
|
|
21
|
+
os.environ["PREDICTIONGUARD_URL"] = "<url>"
|
|
22
|
+
|
|
23
|
+
# Or set your Prediction Guard token and url when initializing the PredictionGuard class.
|
|
24
|
+
client = PredictionGuard(
|
|
25
|
+
api_key="<api_key>",
|
|
26
|
+
url="<url>"
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
result = client.audio.transcriptions.create(
|
|
30
|
+
model="base",
|
|
31
|
+
file="sample_audio.wav"
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
print(json.dumps(
|
|
35
|
+
response,
|
|
36
|
+
sort_keys=True,
|
|
37
|
+
indent=4,
|
|
38
|
+
separators=(",", ": ")
|
|
39
|
+
))
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
def __init__(self, api_key, url, timeout):
|
|
43
|
+
self.api_key = api_key
|
|
44
|
+
self.url = url
|
|
45
|
+
self.timeout = timeout
|
|
46
|
+
|
|
47
|
+
self.transcriptions: AudioTranscriptions = AudioTranscriptions(self.api_key, self.url, self.timeout)
|
|
48
|
+
|
|
49
|
+
class AudioTranscriptions:
|
|
50
|
+
def __init__(self, api_key, url, timeout):
|
|
51
|
+
self.api_key = api_key
|
|
52
|
+
self.url = url
|
|
53
|
+
self.timeout = timeout
|
|
54
|
+
|
|
55
|
+
def create(
|
|
56
|
+
self,
|
|
57
|
+
model: str,
|
|
58
|
+
file: str,
|
|
59
|
+
language: Optional[str] = "auto",
|
|
60
|
+
temperature: Optional[float] = 0.0,
|
|
61
|
+
prompt: Optional[str] = "",
|
|
62
|
+
timestamp_granularities: Optional[List[str]] = None,
|
|
63
|
+
diarization: Optional[bool] = False,
|
|
64
|
+
response_format: Optional[str] = "json",
|
|
65
|
+
toxicity: Optional[bool] = False,
|
|
66
|
+
pii: Optional[str] = "",
|
|
67
|
+
replace_method: Optional[str] = "",
|
|
68
|
+
entity_list: Optional[List[str]] = None,
|
|
69
|
+
injection: Optional[bool] = False,
|
|
70
|
+
) -> Dict[str, Any]:
|
|
71
|
+
"""
|
|
72
|
+
Creates an audio transcription request to the Prediction Guard /audio/transcriptions API
|
|
73
|
+
|
|
74
|
+
:param model: The model to use
|
|
75
|
+
:param file: Audio file to be transcribed
|
|
76
|
+
:param language: The language of the audio file
|
|
77
|
+
:param temperature: The temperature parameter for model transcription
|
|
78
|
+
:param prompt: A prompt to assist in transcription styling
|
|
79
|
+
:param timestamp_granularities: The timestamp granularities to populate for this transcription
|
|
80
|
+
:param diarization: Whether to diarize the audio
|
|
81
|
+
:param response_format: The response format to use
|
|
82
|
+
:param toxicity: Whether to check for output toxicity
|
|
83
|
+
:param pii: Whether to check for or replace pii
|
|
84
|
+
:param replace_method: Replace method for any PII that is present.
|
|
85
|
+
:param entity_list: List of entities to ignore in the PII check.
|
|
86
|
+
:param injection: Whether to check for prompt injection
|
|
87
|
+
:result: A dictionary containing the transcribed text.
|
|
88
|
+
"""
|
|
89
|
+
|
|
90
|
+
# Create a list of tuples, each containing all the parameters for
|
|
91
|
+
# a call to _transcribe_audio
|
|
92
|
+
args = (
|
|
93
|
+
model,
|
|
94
|
+
file,
|
|
95
|
+
language,
|
|
96
|
+
temperature,
|
|
97
|
+
prompt,
|
|
98
|
+
timestamp_granularities,
|
|
99
|
+
diarization,
|
|
100
|
+
response_format,
|
|
101
|
+
pii,
|
|
102
|
+
replace_method,
|
|
103
|
+
entity_list,
|
|
104
|
+
injection,
|
|
105
|
+
toxicity,
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
# Run _transcribe_audio
|
|
109
|
+
choices = self._transcribe_audio(*args)
|
|
110
|
+
return choices
|
|
111
|
+
|
|
112
|
+
def _transcribe_audio(
|
|
113
|
+
self,
|
|
114
|
+
model,
|
|
115
|
+
file,
|
|
116
|
+
language,
|
|
117
|
+
temperature,
|
|
118
|
+
prompt,
|
|
119
|
+
timestamp_granularities,
|
|
120
|
+
diarization,
|
|
121
|
+
response_format,
|
|
122
|
+
pii,
|
|
123
|
+
replace_method,
|
|
124
|
+
entity_list,
|
|
125
|
+
injection,
|
|
126
|
+
toxicity,
|
|
127
|
+
):
|
|
128
|
+
"""
|
|
129
|
+
Function to transcribe an audio file.
|
|
130
|
+
"""
|
|
131
|
+
|
|
132
|
+
headers = {
|
|
133
|
+
"Authorization": "Bearer " + self.api_key,
|
|
134
|
+
"User-Agent": "Prediction Guard Python Client: " + __version__,
|
|
135
|
+
"Toxicity": str(toxicity),
|
|
136
|
+
"Pii": str(pii),
|
|
137
|
+
"Replace-Method": str(replace_method),
|
|
138
|
+
"Entity-List": str(entity_list),
|
|
139
|
+
"Injection": str(injection)
|
|
140
|
+
}
|
|
141
|
+
|
|
142
|
+
if timestamp_granularities:
|
|
143
|
+
if diarization and "segment" in timestamp_granularities:
|
|
144
|
+
raise ValueError(
|
|
145
|
+
"Timestamp granularities cannot be set to "
|
|
146
|
+
"`segments` when using diarization."
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
if response_format != "verbose_json":
|
|
150
|
+
raise ValueError(
|
|
151
|
+
"Response format must be set to `verbose_json` "
|
|
152
|
+
"when using timestamp granularities."
|
|
153
|
+
)
|
|
154
|
+
|
|
155
|
+
if diarization and response_format != "verbose_json":
|
|
156
|
+
raise ValueError(
|
|
157
|
+
"Response format must be set to `verbose_json` "
|
|
158
|
+
"when using diarization."
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
with open(file, "rb") as audio_file:
|
|
162
|
+
files = {"file": (file, audio_file, "audio/wav")}
|
|
163
|
+
data = {
|
|
164
|
+
"model": model,
|
|
165
|
+
"language": language,
|
|
166
|
+
"temperature": temperature,
|
|
167
|
+
"prompt": prompt,
|
|
168
|
+
"timestamp_granularities[]": timestamp_granularities,
|
|
169
|
+
"diarization": str(diarization).lower(),
|
|
170
|
+
"response_format": response_format,
|
|
171
|
+
}
|
|
172
|
+
|
|
173
|
+
response = requests.request(
|
|
174
|
+
"POST", self.url + "/audio/transcriptions", headers=headers, files=files, data=data, timeout=self.timeout
|
|
175
|
+
)
|
|
176
|
+
|
|
177
|
+
# If the request was successful, print the proxies.
|
|
178
|
+
if response.status_code == 200:
|
|
179
|
+
ret = response.json()
|
|
180
|
+
return ret
|
|
181
|
+
elif response.status_code == 429:
|
|
182
|
+
raise ValueError(
|
|
183
|
+
"Could not connect to Prediction Guard API. "
|
|
184
|
+
"Too many requests, rate limit or quota exceeded."
|
|
185
|
+
)
|
|
186
|
+
else:
|
|
187
|
+
# Check if there is a json body in the response. Read that in,
|
|
188
|
+
# print out the error field in the json body, and raise an exception.
|
|
189
|
+
err = ""
|
|
190
|
+
try:
|
|
191
|
+
err = response.json()["error"]
|
|
192
|
+
except Exception:
|
|
193
|
+
pass
|
|
194
|
+
raise ValueError("Could not transcribe the audio file. " + err)
|