Files
doczyai-pipelines/archive/streamlit/ingest.py
T

227 lines
7.6 KiB
Python
Raw Normal View History

2024-02-28 18:23:22 +05:30
import logging
import os
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed
import click
import torch
from langchain.docstore.document import Document
from langchain.embeddings import HuggingFaceInstructEmbeddings
from langchain.text_splitter import Language, RecursiveCharacterTextSplitter
from langchain.vectorstores import Chroma
import uuid
from constants import (
CHROMA_SETTINGS,
DOCUMENT_MAP,
EMBEDDING_MODEL_NAME,
INGEST_THREADS,
PERSIST_DIRECTORY,
SOURCE_DIRECTORY,
)
import boto3
from langchain.embeddings.bedrock import BedrockEmbeddings
2024-09-30 12:21:29 +01:00
2024-02-28 18:23:22 +05:30
def file_log(logentry):
2024-09-30 12:21:29 +01:00
file1 = open("file_ingest.log", "a")
file1.write(logentry + "\n")
file1.close()
print(logentry + "\n")
2024-02-28 18:23:22 +05:30
def load_single_document(file_path: str) -> Document:
# Loads a single document from a file path
try:
2024-09-30 12:21:29 +01:00
file_extension = os.path.splitext(file_path)[1]
loader_class = DOCUMENT_MAP.get(file_extension)
if loader_class:
file_log(file_path + " loaded.")
loader = loader_class(file_path)
else:
file_log(file_path + " document type is undefined.")
raise ValueError("Document type is undefined")
return loader.load()[0]
2024-02-28 18:23:22 +05:30
except Exception as ex:
2024-09-30 12:21:29 +01:00
file_log("%s loading error: \n%s" % (file_path, ex))
return None
2024-02-28 18:23:22 +05:30
def load_document_batch(filepaths):
logging.info("Loading document batch")
# create a thread pool
with ThreadPoolExecutor(len(filepaths)) as exe:
# load files
futures = [exe.submit(load_single_document, name) for name in filepaths]
# collect data
if futures is None:
2024-09-30 12:21:29 +01:00
file_log(name + " failed to submit")
return None
2024-02-28 18:23:22 +05:30
else:
2024-09-30 12:21:29 +01:00
data_list = [future.result() for future in futures]
# return data and file paths
return (data_list, filepaths)
2024-02-28 18:23:22 +05:30
def load_documents(source_dir: str) -> list[Document]:
# Loads all documents from the source documents directory, including nested folders
paths = []
for root, _, files in os.walk(source_dir):
for file_name in files:
2024-09-30 12:21:29 +01:00
print("Importing: " + file_name)
2024-02-28 18:23:22 +05:30
file_extension = os.path.splitext(file_name)[1]
source_file_path = os.path.join(root, file_name)
if file_extension in DOCUMENT_MAP.keys():
paths.append(source_file_path)
# Have at least one worker and at most INGEST_THREADS workers
n_workers = min(INGEST_THREADS, max(len(paths), 1))
chunksize = round(len(paths) / n_workers)
docs = []
with ProcessPoolExecutor(n_workers) as executor:
futures = []
# split the load operations into chunks
for i in range(0, len(paths), chunksize):
# select a chunk of filenames
filepaths = paths[i : (i + chunksize)]
# submit the task
try:
2024-09-30 12:21:29 +01:00
future = executor.submit(load_document_batch, filepaths)
2024-02-28 18:23:22 +05:30
except Exception as ex:
2024-09-30 12:21:29 +01:00
file_log("executor task failed: %s" % (ex))
future = None
2024-02-28 18:23:22 +05:30
if future is not None:
2024-09-30 12:21:29 +01:00
futures.append(future)
2024-02-28 18:23:22 +05:30
# process all results
for future in as_completed(futures):
# open the file and load the data
try:
contents, _ = future.result()
docs.extend(contents)
except Exception as ex:
2024-09-30 12:21:29 +01:00
file_log("Exception: %s" % (ex))
2024-02-28 18:23:22 +05:30
return docs
def split_documents(documents: list[Document]) -> tuple[list[Document], list[Document]]:
# Splits documents for correct Text Splitter
text_docs, python_docs = [], []
for doc in documents:
if doc is not None:
2024-09-30 12:21:29 +01:00
file_extension = os.path.splitext(doc.metadata["source"])[1]
if file_extension == ".py":
python_docs.append(doc)
else:
text_docs.append(doc)
2024-02-28 18:23:22 +05:30
return text_docs, python_docs
2024-09-30 12:21:29 +01:00
2024-02-28 18:23:22 +05:30
def process_in_batches(texts, batch_size):
for i in range(0, len(texts), batch_size):
2024-09-30 12:21:29 +01:00
yield texts[i : i + batch_size]
2024-02-28 18:23:22 +05:30
@click.command()
@click.option(
"--device_type",
default="cuda" if torch.cuda.is_available() else "cpu",
type=click.Choice(
[
"cpu",
"cuda",
"ipu",
"xpu",
"mkldnn",
"opengl",
"opencl",
"ideep",
"hip",
"ve",
"fpga",
"ort",
"xla",
"lazy",
"vulkan",
"mps",
"meta",
"hpu",
"mtia",
],
),
help="Device to run on. (Default is cuda)",
)
def main(device_type):
# Load documents and split in chunks
logging.info(f"Loading documents from {SOURCE_DIRECTORY}")
2024-09-30 12:21:29 +01:00
for filename in os.listdir(SOURCE_DIRECTORY):
with open(os.path.join(SOURCE_DIRECTORY, filename), "r") as infile:
2024-02-28 18:23:22 +05:30
text = infile.read()
2024-09-30 12:21:29 +01:00
text_splitted = [i for i in text.split("Start of Page No. = ")]
2024-02-28 18:23:22 +05:30
for i, txt in enumerate(text_splitted):
if len(txt) > 2:
2024-09-30 12:21:29 +01:00
page_path = os.path.join(
SOURCE_DIRECTORY, f"{filename[:-4]}_page{i}.txt"
)
with open(page_path, "w") as f:
2024-02-28 18:23:22 +05:30
f.write(txt)
documents = [load_single_document(page_path)]
os.remove(page_path)
text_documents, python_documents = split_documents(documents)
2024-09-30 12:21:29 +01:00
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=1000, chunk_overlap=200
)
2024-02-28 18:23:22 +05:30
python_splitter = RecursiveCharacterTextSplitter.from_language(
language=Language.PYTHON, chunk_size=880, chunk_overlap=200
)
texts = text_splitter.split_documents(text_documents)
texts.extend(python_splitter.split_documents(python_documents))
2024-09-30 12:21:29 +01:00
logging.info(
f"Loaded {len(documents)} documents from {SOURCE_DIRECTORY}"
)
2024-02-28 18:23:22 +05:30
logging.info(f"Split into {len(texts)} chunks of text")
# Create embeddings
# embeddings = HuggingFaceInstructEmbeddings(
# model_name=EMBEDDING_MODEL_NAME,
# model_kwargs={"device": device_type},
# )
bedrock_runtime = boto3.client(
service_name="bedrock-runtime",
region_name="us-east-1",
)
embeddings = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
"""
db = Chroma.from_documents(
texts,
embeddings,
persist_directory=PERSIST_DIRECTORY,
client_settings=CHROMA_SETTINGS,
)
"""
# for batch_texts in process_in_batches(texts, 20000): # https://github.com/PromtEngineer/localGPT/issues/489y
print(filename)
db = Chroma.from_documents(
texts,
embeddings,
persist_directory=PERSIST_DIRECTORY,
client_settings=CHROMA_SETTINGS,
collection_metadata={"hnsw:space": "cosine"},
# ids = [str(filename)]
)
if __name__ == "__main__":
logging.basicConfig(
2024-09-30 12:21:29 +01:00
format="%(asctime)s - %(levelname)s - %(filename)s:%(lineno)s - %(message)s",
level=logging.INFO,
2024-02-28 18:23:22 +05:30
)
main()