Fix doc RAG processing, put graph-rag and doc-rag in appropriate component

files.
This commit is contained in:
Cyber MacGeddon 2025-01-04 21:35:22 +00:00
parent 906417c54c
commit 6c3cd607d1
11 changed files with 245 additions and 180 deletions

View file

@ -39,5 +39,35 @@ local prompts = import "prompts/mixtral.jsonnet";
}, },
"document-embeddings" +: {
create:: function(engine)
local container =
engine.container("document-embeddings")
.with_image(images.trustgraph)
.with_command([
"document-embeddings",
"-p",
url.pulsar,
])
.with_limits("1.0", "512M")
.with_reservations("0.5", "512M");
local containerSet = engine.containers(
"document-embeddings", [ container ]
);
local service =
engine.internalService(containerSet)
.with_port(8000, 8000, "metrics");
engine.resources([
containerSet,
service,
])
},
} }

View file

@ -138,5 +138,35 @@ local url = import "values/url.jsonnet";
}, },
"graph-embeddings" +: {
create:: function(engine)
local container =
engine.container("graph-embeddings")
.with_image(images.trustgraph)
.with_command([
"graph-embeddings",
"-p",
url.pulsar,
])
.with_limits("1.0", "512M")
.with_reservations("0.5", "512M");
local containerSet = engine.containers(
"graph-embeddings", [ container ]
);
local service =
engine.internalService(containerSet)
.with_port(8000, 8000, "metrics");
engine.resources([
containerSet,
service,
])
},
} }

View file

@ -119,66 +119,6 @@ local prompt = import "prompt-template.jsonnet";
}, },
"graph-embeddings" +: {
create:: function(engine)
local container =
engine.container("graph-embeddings")
.with_image(images.trustgraph)
.with_command([
"graph-embeddings",
"-p",
url.pulsar,
])
.with_limits("1.0", "512M")
.with_reservations("0.5", "512M");
local containerSet = engine.containers(
"graph-embeddings", [ container ]
);
local service =
engine.internalService(containerSet)
.with_port(8000, 8000, "metrics");
engine.resources([
containerSet,
service,
])
},
"document-embeddings" +: {
create:: function(engine)
local container =
engine.container("document-embeddings")
.with_image(images.trustgraph)
.with_command([
"document-embeddings",
"-p",
url.pulsar,
])
.with_limits("1.0", "512M")
.with_reservations("0.5", "512M");
local containerSet = engine.containers(
"document-embeddings", [ container ]
);
local service =
engine.internalService(containerSet)
.with_port(8000, 8000, "metrics");
engine.resources([
containerSet,
service,
])
},
"metering" +: { "metering" +: {
create:: function(engine) create:: function(engine)

View file

@ -131,6 +131,35 @@ class Api:
except: except:
raise ProtocolException(f"Response not formatted correctly") raise ProtocolException(f"Response not formatted correctly")
def document_rag(self, question):
# The input consists of a question
input = {
"query": question
}
url = f"{self.url}document-rag"
# Invoke the API, input is passed as JSON
resp = requests.post(url, json=input)
# Should be a 200 status code
if resp.status_code != 200:
raise ProtocolException(f"Status code {resp.status_code}")
try:
# Parse the response as JSON
object = resp.json()
except:
raise ProtocolException(f"Expected JSON response")
self.check_error(resp)
try:
return object["response"]
except:
raise ProtocolException(f"Response not formatted correctly")
def embeddings(self, text): def embeddings(self, text):
# The input consists of a text block # The input consists of a text block

View file

@ -38,8 +38,12 @@ class DocumentEmbeddingsClient(BaseClient):
output_schema=DocumentEmbeddingsResponse, output_schema=DocumentEmbeddingsResponse,
) )
def request(self, vectors, limit=10, timeout=300): def request(
self, vectors, user="trustgraph", collection="default",
limit=10, timeout=300
):
return self.call( return self.call(
user=user, collection=collection,
vectors=vectors, limit=limit, timeout=timeout vectors=vectors, limit=limit, timeout=timeout
).documents ).documents

View file

@ -55,6 +55,8 @@ document_embeddings_store_queue = topic('document-embeddings-store')
class DocumentEmbeddingsRequest(Record): class DocumentEmbeddingsRequest(Record):
vectors = Array(Array(Double())) vectors = Array(Array(Double()))
limit = Integer() limit = Integer()
user = String()
collection = String()
class DocumentEmbeddingsResponse(Record): class DocumentEmbeddingsResponse(Record):
error = Error() error = Error()

View file

@ -16,6 +16,44 @@ from . schema import document_embeddings_response_queue
LABEL="http://www.w3.org/2000/01/rdf-schema#label" LABEL="http://www.w3.org/2000/01/rdf-schema#label"
DEFINITION="http://www.w3.org/2004/02/skos/core#definition" DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query:
def __init__(self, rag, user, collection, verbose):
self.rag = rag
self.user = user
self.collection = collection
self.verbose = verbose
def get_vector(self, query):
if self.verbose:
print("Compute embeddings...", flush=True)
qembeds = self.rag.embeddings.request(query)
if self.verbose:
print("Done.", flush=True)
return qembeds
def get_docs(self, query):
vectors = self.get_vector(query)
if self.verbose:
print("Get entities...", flush=True)
docs = self.rag.de_client.request(
vectors, limit=self.rag.doc_limit
)
if self.verbose:
print("Docs:", flush=True)
for doc in docs:
print(doc, flush=True)
return docs
class DocumentRag: class DocumentRag:
def __init__( def __init__(
@ -55,7 +93,7 @@ class DocumentRag:
print("Initialising...", flush=True) print("Initialising...", flush=True)
# FIXME: Configurable # FIXME: Configurable
self.entity_limit = 20 self.doc_limit = 20
self.de_client = DocumentEmbeddingsClient( self.de_client = DocumentEmbeddingsClient(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
@ -81,42 +119,16 @@ class DocumentRag:
if self.verbose: if self.verbose:
print("Initialised", flush=True) print("Initialised", flush=True)
def get_vector(self, query): def query(self, query, user="trustgraph", collection="default"):
if self.verbose:
print("Compute embeddings...", flush=True)
qembeds = self.embeddings.request(query)
if self.verbose:
print("Done.", flush=True)
return qembeds
def get_docs(self, query):
vectors = self.get_vector(query)
if self.verbose:
print("Get entities...", flush=True)
docs = self.de_client.request(
vectors, self.entity_limit
)
if self.verbose:
print("Docs:", flush=True)
for doc in docs:
print(doc, flush=True)
return docs
def query(self, query):
if self.verbose: if self.verbose:
print("Construct prompt...", flush=True) print("Construct prompt...", flush=True)
docs = self.get_docs(query) q = Query(
rag=self, user=user, collection=collection, verbose=self.verbose
)
docs = q.get_docs(query)
if self.verbose: if self.verbose:
print("Invoke LLM...", flush=True) print("Invoke LLM...", flush=True)

View file

@ -31,6 +31,7 @@ from . subscriber import Subscriber
from . text_completion import TextCompletionRequestor from . text_completion import TextCompletionRequestor
from . prompt import PromptRequestor from . prompt import PromptRequestor
from . graph_rag import GraphRagRequestor from . graph_rag import GraphRagRequestor
from . document_rag import DocumentRagRequestor
from . triples_query import TriplesQueryRequestor from . triples_query import TriplesQueryRequestor
from . graph_embeddings_query import GraphEmbeddingsQueryRequestor from . graph_embeddings_query import GraphEmbeddingsQueryRequestor
from . embeddings import EmbeddingsRequestor from . embeddings import EmbeddingsRequestor
@ -91,6 +92,10 @@ class Api:
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth,
), ),
"document-rag": DocumentRagRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth,
),
"triples-query": TriplesQueryRequestor( "triples-query": TriplesQueryRequestor(
pulsar_host=self.pulsar_host, timeout=self.timeout, pulsar_host=self.pulsar_host, timeout=self.timeout,
auth = self.auth, auth = self.auth,
@ -140,6 +145,10 @@ class Api:
endpoint_path = "/api/v1/graph-rag", auth=self.auth, endpoint_path = "/api/v1/graph-rag", auth=self.auth,
requestor = self.services["graph-rag"], requestor = self.services["graph-rag"],
), ),
ServiceEndpoint(
endpoint_path = "/api/v1/document-rag", auth=self.auth,
requestor = self.services["document-rag"],
),
ServiceEndpoint( ServiceEndpoint(
endpoint_path = "/api/v1/triples-query", auth=self.auth, endpoint_path = "/api/v1/triples-query", auth=self.auth,
requestor = self.services["triples-query"], requestor = self.services["triples-query"],

View file

@ -39,9 +39,14 @@ class Processor(Consumer):
v = msg.value() v = msg.value()
chunk = v.chunk.decode("utf-8") for emb in v.chunks:
if v.chunk != "" and v.chunk is not None: chunk = emb.chunk.decode("utf-8")
if chunk == "" or chunk is None: continue
for vec in emb.vectors:
if chunk != "" and v.chunk is not None:
for vec in v.vectors: for vec in v.vectors:
self.vecstore.insert(vec, chunk) self.vecstore.insert(vec, chunk)

View file

@ -65,9 +65,12 @@ class Processor(Consumer):
v = msg.value() v = msg.value()
chunk = v.chunk.decode("utf-8") for emb in v.chunks:
if chunk == "": return chunk = emb.chunk.decode("utf-8")
if chunk == "" or chunk is None: continue
for vec in emb.vectors:
for vec in v.vectors: for vec in v.vectors:

View file

@ -44,11 +44,12 @@ class Processor(Consumer):
v = msg.value() v = msg.value()
chunk = v.chunk.decode("utf-8") for emb in v.chunks:
chunk = emb.chunk.decode("utf-8")
if chunk == "": return if chunk == "": return
for vec in v.vectors: for vec in emb.vectors:
dim = len(vec) dim = len(vec)
collection = ( collection = (