LangChain / app /api /api_v1 /endpoints /openai.py
dcd018's picture
Labelstudio config.
98fa7bf
Raw
History Blame Contribute Delete
10.2 kB
import uuid
from typing import List, Tuple, Any
from app import schemas, label_studio_project
from app.core.config import settings
from app.api import deps
from app.callbacks.chain_stream import chat
from app.vectorstores import PGVector
from sqlalchemy.orm import Session
from fastapi import APIRouter, Depends, BackgroundTasks
from fastapi.responses import StreamingResponse
from langchain.document_loaders import PyPDFLoader
from langchain.embeddings.openai import OpenAIEmbeddings
from langchain.document_loaders import TextLoader
from langchain.docstore.document import Document
import requests
from pdf2image import convert_from_bytes
import PIL
import os
import json
import logging
from sqlalchemy.sql.expression import cast
from sqlalchemy.dialects.postgresql import UUID, JSON
import zipfile
import glob
import papermill as pm
import shutil
import nbformat as nbf
import random
import time
import requests
from app.core.config import settings
import pathlib
import io
from langchain.vectorstores.pgvector import (
CollectionStore,
DistanceStrategy
)
from langchain.schema import (
HumanMessage,
SystemMessage
)
PIL.Image.MAX_IMAGE_PIXELS = None
router = APIRouter()
embeddings = OpenAIEmbeddings()
def retry_with_backoff(retries = 5, backoff_in_seconds = 1):
def rwb(f):
def wrapper(*args, **kwargs):
x = 0
while True:
try:
return f(*args, **kwargs)
except:
if x == retries:
raise
sleep = (backoff_in_seconds * 2 ** x +
random.uniform(0, 1))
time.sleep(sleep)
x += 1
return wrapper
return rwb
def convert_from_ls(result):
if 'original_width' not in result or 'original_height' not in result:
return None
value = result['value']
w, h = result['original_width'], result['original_height']
if all([key in value for key in ['x', 'y', 'width', 'height']]):
return w * value['x'] / 100.0, \
h * value['y'] / 100.0, \
w * value['width'] / 100.0, \
h * value['height'] / 100.0
def embed_documents(collection: CollectionStore):
loader: TextLoader = PyPDFLoader(collection.cmetadata["location"])
documents: List[Document] = loader.load_and_split()
PGVector.from_documents(
embedding=embeddings,
documents=documents,
collection_name=collection.name,
connection_string=settings.SQLALCHEMY_DATABASE_URI,
pre_delete_embeddings=True
)
@retry_with_backoff(retries=6)
def import_task(tmp_file):
params = {'return_task_ids': '1'}
with open(tmp_file, mode='rb') as f:
response = requests.post(
"{}/api/projects/{}/import".format(settings.LABEL_STUDIO_URL, label_studio_project.id),
headers={"Authorization": "Token {}".format(settings.LABEL_STUDIO_API_KEY)},
timeout=(600, 600),
params=params,
files={'file': f})
response = response.json()
if 'import' in response:
# check import status
timeout = 500
fibonacci_backoff = [1, 1]
start_time = time.time()
while True:
import_status = requests.post(
"{}/api/projects/{}/imports/{}".format(settings.LABEL_STUDIO_URL, label_studio_project.id, response["import"]),
headers={"Authorization": "Token {}".format(settings.LABEL_STUDIO_API_KEY)},
params=params,
timeout=(600, 600)).json()
if import_status['status'] == 'completed':
return import_status['task_ids']
if import_status['status'] == 'failed':
raise Exception(import_status['error'])
if time.time() - start_time >= timeout:
raise Exception('Import timeout')
time.sleep(fibonacci_backoff[0])
fibonacci_backoff = [
fibonacci_backoff[1],
fibonacci_backoff[0] + fibonacci_backoff[1],
]
print(response)
return response['task_ids']
@retry_with_backoff(retries=6)
def export_task(task_id, tmp_archive):
response = requests.get(
"{}/api/projects/{}/export?export_type=COCO&ids[]={}".format(settings.LABEL_STUDIO_URL, label_studio_project.id, task_id),
headers={"Authorization": "Token {}".format(settings.LABEL_STUDIO_API_KEY)},
timeout=(500, 500),
stream=True)
export_path = pathlib.Path(tmp_archive)
export_path.parent.mkdir(parents=True, exist_ok=True)
with open(export_path, "wb") as out_file:
for chunk in response.iter_content(chunk_size=62):
out_file.write(chunk)
@retry_with_backoff(retries=6)
def embed(db: Session, payload: schemas.LabelStudio):
task_id = payload.task["id"]
img_name = os.path.basename(payload.task["data"]["ocr"])
annotations = [convert_from_ls(result) for result in payload.annotation["result"] if result["type"] == "rectangle"]
collection = db.query(CollectionStore).filter(CollectionStore.cmetadata["task_id"].astext == str(task_id)).first()
tmp_dir = "/tmp/{}".format(collection.uuid)
tmp_archive = "{}.zip".format(tmp_dir)
section_dir = "{}/sections".format(tmp_dir)
notebook_dir = "{}/notebooks".format(tmp_dir)
export_task(task_id, tmp_archive)
with zipfile.ZipFile(tmp_archive, 'r') as zip_ref:
zip_ref.extractall(tmp_dir)
try:
os.makedirs(section_dir)
img_path = "{}/images/{}".format(tmp_dir, img_name)
with PIL.Image.open(img_path) as image:
for i, annotation in enumerate(annotations):
x, y, width, height = annotation
cropped = image.crop([
x,
y,
width + x,
height + y
])
cropped.save("{}/section-{}.jpg".format(section_dir, i))
section_paths = glob.glob("{}/*".format(section_dir))
os.makedirs(notebook_dir)
for embedding in collection.embeddings:
db.delete(embedding)
db.commit()
db.refresh(collection)
notebooks = []
for path in section_paths:
section_path = "{}/{}.ipynb".format(notebook_dir, os.path.basename(path))
pm.execute_notebook(
'/home/appuser/app/notebooks/inference.ipynb',
section_path,
parameters=dict(file_path=path)
)
notebook = nbf.read(section_path, nbf.NO_CONVERT)
notebooks.append(nbf.writes(notebook, nbf.NO_CONVERT))
cell = [cell for cell in notebook.cells if "output" in cell.metadata.tags][0]
texts = [output.text for output in cell.outputs]
PGVector.from_texts(
embedding=embeddings,
texts=texts,
collection_name=collection.name,
connection_string=settings.SQLALCHEMY_DATABASE_URI,
pre_delete_embeddings=True,
pre_delete_collection=False
)
shutil.rmtree(tmp_dir)
print("Done.")
except:
shutil.rmtree(tmp_dir)
raise
@router.post("/embeddings", response_model=List[float])
async def create_embedding(text: str) -> List[float]:
"""
Create embedding.
"""
return embeddings.embed_query(text)
@router.post("/collections", response_model=schemas.CollectionStoreInDB)
async def create_collection(
collection_in: schemas.CollectionStoreCreate,
background_tasks: BackgroundTasks,
db: Session = Depends(deps.get_db),
) -> CollectionStore:
"""
Create a collection of embeddings.
"""
collection_name = str(uuid.uuid4())
tmp_file = "/tmp/{}.jpg".format(collection_name)
results = requests.get(collection_in.location, timeout=(500, 500))
pages = convert_from_bytes(results.content, 500, None, 1, 1)
cover = pages[0]
width, height = cover.size
target_width = round(width - (width * 0.8))
target_height = round(height - (height * 0.8))
cover = cover.resize((target_width, target_height), PIL.Image.Resampling.LANCZOS)
cover.save(tmp_file, optimize=True, quality=75)
tasks = import_task(tmp_file)
os.remove(tmp_file)
collection_in.metadata["location"] = collection_in.location
collection_in.metadata["task_id"] = tasks[0]
collection, created = CollectionStore.get_or_create(
db, collection_name, cmetadata=collection_in.metadata
)
#background_tasks.add_task(embed_documents, collection)
return collection
@router.get("/ask")
async def ask_question(
question: str,
collection_uuid: str = None,
db: Session = Depends(deps.get_db),
) -> Any:
"""
Ask a question.
"""
messages = [
HumanMessage(content=question)
]
if collection_uuid:
collection_name: CollectionStore = (db.query(CollectionStore.name)
.filter_by(uuid = uuid.UUID(collection_uuid))
.scalar())
if collection_name is not None:
store: PGVector = PGVector(
connection_string=settings.SQLALCHEMY_DATABASE_URI,
embedding_function=embeddings,
collection_name=collection_name,
distance_strategy=DistanceStrategy.COSINE
)
documents: List[Tuple[Document, float]] = store.similarity_search_with_score(query=question, k=2)
for document in documents:
(document, score) = document
content: str = document.page_content
messages.append(SystemMessage(content=content))
return StreamingResponse(chat(messages), media_type='text/event-stream')
@router.post("/inference")
async def inference(
payload: schemas.LabelStudio,
background_tasks: BackgroundTasks,
db: Session = Depends(deps.get_db)
):
"""
Create embedding.
"""
embed(db, payload)
#background_tasks.add_task(embed, db, payload)
return { "success": True }