lawliet/CS224-knowledge-discovery
1
1import os2from typing import Union, Any, Optional, Mapping, cast3from collections.abc import Callable4import aiohttp5import asyncio6 7 8async def transform_json(resp):9 return await resp.json()10 11 12async def make_async_request(13 *,14 uri: str,15 method: str = "GET",16 headers=None,17 params=None,18 total_retries: Optional[int] = None,19 transform=transform_json,20 json=None,21):22 """Attempts to parse json by default"""23 total_retries = 6 if total_retries is None else total_retries24 retry_count = 025 sleep_time = pow(2, retry_count)26 27 while retry_count <= total_retries:28 if retry_count > 0:29 print(f"Retry {retry_count}/{total_retries}, sleeping {sleep_time}s")30 await asyncio.sleep(sleep_time)31 sleep_time = pow(2, retry_count) # set sleep time for next cycle32 async with aiohttp.ClientSession() as session, session.request(33 method, uri, headers=headers, params=params, json=json34 ) as response:35 if 200 <= response.status < 300:36 return await transform(response)37 38 retry_count += 139 raise Exception("Request Failed")40 41 42class TextEncoder:43 def __init__(self, model_name: str = "text-embedding-ada-002") -> None:44 super().__init__()45 self.OPENAI_API_KEY = os.environ["OPENAI_API_KEY"]46 self.embedding_endpoint = "https://api.openai.com/v1/embeddings"47 self.model_name = model_name48 49 async def encode_text(self, texts):50 response = await make_async_request(51 uri=self.embedding_endpoint,52 method="POST",53 headers={54 "content_type": "application/json",55 "Authorization": f"Bearer {self.OPENAI_API_KEY}",56 },57 json={58 "input": texts,59 "model": self.model_name,60 },61 total_retries=2,62 )63 # get embedding, OpenAI API doesn't ensure the output would be sorted64 return [65 data["embedding"]66 for data in sorted(response["data"], key=lambda x: x["index"])67 ], response["usage"]["total_tokens"]68 