Change document-rag and graph-rag processing so that the user can

specify parameters.  Changes in Pulsar services, Pulsar message
schemas, gateway and command-line tools.  User-visible changes in
new parameters on command-line tools.
This commit is contained in:
Cyber MacGeddon 2025-03-12 17:51:47 +00:00
parent f1559c5944
commit 8830e99e57
12 changed files with 154 additions and 48 deletions

View file

@ -5,6 +5,8 @@ local prompts = import "prompts/mixtral.jsonnet";
{ {
"document-rag-doc-limit":: 20,
"document-rag" +: { "document-rag" +: {
create:: function(engine) create:: function(engine)

View file

@ -4,9 +4,9 @@ local url = import "values/url.jsonnet";
{ {
"graph-rag-entity-limit":: 50, "graph-rag-entity-limit":: 20,
"graph-rag-triple-limit":: 30, "graph-rag-triple-limit":: 10,
"graph-rag-max-subgraph-size":: 3000, "graph-rag-max-subgraph-size":: 1000,
"kg-extract-definitions" +: { "kg-extract-definitions" +: {

View file

@ -102,11 +102,19 @@ class Api:
except: except:
raise ProtocolException(f"Response not formatted correctly") raise ProtocolException(f"Response not formatted correctly")
def graph_rag(self, question): def graph_rag(
self, question, user="trustgraph", collection="default",
entity_limit=50, triple_limit=30, subgraph_limit=1000,
):
# The input consists of a question # The input consists of a question
input = { input = {
"query": question "query": question,
"user": user,
"collection": collection,
"entity-limit": entity_limit,
"triple-limit": triple_limit,
"max-subgraph-limit": subgraph_limit,
} }
url = f"{self.url}graph-rag" url = f"{self.url}graph-rag"
@ -131,11 +139,17 @@ class Api:
except: except:
raise ProtocolException(f"Response not formatted correctly") raise ProtocolException(f"Response not formatted correctly")
def document_rag(self, question): def document_rag(
self, question, user="trustgraph", collection="default",
doc_limit=10,
):
# The input consists of a question # The input consists of a question
input = { input = {
"query": question "query": question,
"user": user,
"collection": collection,
"doc-limit": doc_limit,
} }
url = f"{self.url}document-rag" url = f"{self.url}document-rag"

View file

@ -11,6 +11,9 @@ class GraphRagQuery(Record):
query = String() query = String()
user = String() user = String()
collection = String() collection = String()
entity_limit = Integer()
triple_limit = Integer()
max_subgraph_size = Integer()
class GraphRagResponse(Record): class GraphRagResponse(Record):
error = Error() error = Error()
@ -31,6 +34,7 @@ class DocumentRagQuery(Record):
query = String() query = String()
user = String() user = String()
collection = String() collection = String()
doc_limit = Integer()
class DocumentRagResponse(Record): class DocumentRagResponse(Record):
error = Error() error = Error()

View file

@ -11,13 +11,16 @@ from trustgraph.api import Api
default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/') default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
default_doc_limit = 10
def question(url, question, user, collection): def question(url, question, user, collection, doc_limit):
rag = Api(url) rag = Api(url)
# user=user, collection=collection, resp = rag.document_rag(
resp = rag.document_rag(question=question) question=question, user=user, collection=collection,
doc_limit=doc_limit,
)
print(resp) print(resp)
@ -58,6 +61,12 @@ def main():
help=f'Collection ID (default: {default_collection})' help=f'Collection ID (default: {default_collection})'
) )
parser.add_argument(
'-d', '--doc-limit',
default=default_doc_limit,
help=f'Document limit (default: {default_doc_limit})'
)
args = parser.parse_args() args = parser.parse_args()
try: try:
@ -67,6 +76,7 @@ def main():
question=args.question, question=args.question,
user=args.user, user=args.user,
collection=args.collection, collection=args.collection,
doc_limit=args.doc_limit,
) )
except Exception as e: except Exception as e:

View file

@ -12,12 +12,18 @@ default_url = os.getenv("TRUSTGRAPH_URL", 'http://localhost:8088/')
default_user = 'trustgraph' default_user = 'trustgraph'
default_collection = 'default' default_collection = 'default'
def question(url, question, user, collection): def question(
url, question, user, collection, entity_limit, triple_limit,
subgraph_limit
):
rag = Api(url) rag = Api(url)
# user=user, collection=collection, resp = rag.graph_rag(
resp = rag.graph_rag(question=question) question=question, user=user, collection=collection,
entity_limit=entity_limit, triple_limit=triple_limit,
subgraph_limit=subgraph_limit
)
print(resp) print(resp)
@ -52,6 +58,24 @@ def main():
help=f'Collection ID (default: {default_collection})' help=f'Collection ID (default: {default_collection})'
) )
parser.add_argument(
'-e', '--entity-limit',
default=default_entity_limit,
help=f'Entity limit (default: {default_entity_limit})'
)
parser.add_argument(
'-t', '--triple-limit',
default=default_triple_limit,
help=f'Triple limit (default: {default_triple_limit})'
)
parser.add_argument(
'-s', '--subgraph-limit',
default=default_subgraph_limit,
help=f'Subgraph limit (default: {subgraph_limit})'
)
args = parser.parse_args() args = parser.parse_args()
try: try:
@ -61,6 +85,9 @@ def main():
question=args.question, question=args.question,
user=args.user, user=args.user,
collection=args.collection, collection=args.collection,
entity_limit=args.entity_limit,
triple_limit=args.triple_limit,
subgraph_limit=args.subgraph_limit,
) )
except Exception as e: except Exception as e:

View file

@ -18,11 +18,15 @@ DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query: class Query:
def __init__(self, rag, user, collection, verbose): def __init__(
self, rag, user, collection, verbose,
doc_limit=20
):
self.rag = rag self.rag = rag
self.user = user self.user = user
self.collection = collection self.collection = collection
self.verbose = verbose self.verbose = verbose
self.doc_limit = doc_limit
def get_vector(self, query): def get_vector(self, query):
@ -44,7 +48,7 @@ class Query:
print("Get entities...", flush=True) print("Get entities...", flush=True)
docs = self.rag.de_client.request( docs = self.rag.de_client.request(
vectors, limit=self.rag.doc_limit vectors, limit=self.doc_limit
) )
if self.verbose: if self.verbose:
@ -93,9 +97,6 @@ class DocumentRag:
if self.verbose: if self.verbose:
print("Initialising...", flush=True) print("Initialising...", flush=True)
# FIXME: Configurable
self.doc_limit = 20
self.de_client = DocumentEmbeddingsClient( self.de_client = DocumentEmbeddingsClient(
pulsar_host=pulsar_host, pulsar_host=pulsar_host,
subscriber=module + "-de", subscriber=module + "-de",
@ -123,13 +124,17 @@ class DocumentRag:
if self.verbose: if self.verbose:
print("Initialised", flush=True) print("Initialised", flush=True)
def query(self, query, user="trustgraph", collection="default"): def query(
self, query, user="trustgraph", collection="default",
doc_limit=20,
):
if self.verbose: if self.verbose:
print("Construct prompt...", flush=True) print("Construct prompt...", flush=True)
q = Query( q = Query(
rag=self, user=user, collection=collection, verbose=self.verbose rag=self, user=user, collection=collection, verbose=self.verbose,
doc_limit=doc_limit
) )
docs = q.get_docs(query) docs = q.get_docs(query)

View file

@ -23,6 +23,7 @@ class DocumentRagRequestor(ServiceRequestor):
query=body["query"], query=body["query"],
user=body.get("user", "trustgraph"), user=body.get("user", "trustgraph"),
collection=body.get("collection", "default"), collection=body.get("collection", "default"),
doc_limit=body.get("doc-limit", 20),
) )
def from_response(self, message): def from_response(self, message):

View file

@ -23,6 +23,9 @@ class GraphRagRequestor(ServiceRequestor):
query=body["query"], query=body["query"],
user=body.get("user", "trustgraph"), user=body.get("user", "trustgraph"),
collection=body.get("collection", "default"), collection=body.get("collection", "default"),
entity_limit=body.get("entity-limit", 50),
triple_limit=body.get("triple-limit", 30),
doc_limit=body.get("doc-limit", 1000),
) )
def from_response(self, message): def from_response(self, message):

View file

@ -20,11 +20,17 @@ DEFINITION="http://www.w3.org/2004/02/skos/core#definition"
class Query: class Query:
def __init__(self, rag, user, collection, verbose): def __init__(
self, rag, user, collection, verbose,
entity_limit=50, query_limit=30, max_subgraph_size=1000,
):
self.rag = rag self.rag = rag
self.user = user self.user = user
self.collection = collection self.collection = collection
self.verbose = verbose self.verbose = verbose
self.entity_limit = entity_limit
self.query_limit = query_limit
self.max_subgraph_size = max_subgraph_size
def get_vector(self, query): def get_vector(self, query):
@ -47,7 +53,7 @@ class Query:
entities = self.rag.ge_client.request( entities = self.rag.ge_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
vectors=vectors, limit=self.rag.entity_limit, vectors=vectors, limit=self.entity_limit,
) )
entities = [ entities = [
@ -93,7 +99,7 @@ class Query:
res = self.rag.triples_client.request( res = self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=e, p=None, o=None, s=e, p=None, o=None,
limit=self.rag.query_limit limit=self.query_limit
) )
for triple in res: for triple in res:
@ -104,7 +110,7 @@ class Query:
res = self.rag.triples_client.request( res = self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=None, p=e, o=None, s=None, p=e, o=None,
limit=self.rag.query_limit limit=self.query_limit
) )
for triple in res: for triple in res:
@ -115,7 +121,7 @@ class Query:
res = self.rag.triples_client.request( res = self.rag.triples_client.request(
user=self.user, collection=self.collection, user=self.user, collection=self.collection,
s=None, p=None, o=e, s=None, p=None, o=e,
limit=self.rag.query_limit, limit=self.query_limit,
) )
for triple in res: for triple in res:
@ -125,7 +131,7 @@ class Query:
subgraph = list(subgraph) subgraph = list(subgraph)
subgraph = subgraph[0:self.rag.max_subgraph_size] subgraph = subgraph[0:self.max_subgraph_size]
if self.verbose: if self.verbose:
print("Subgraph:", flush=True) print("Subgraph:", flush=True)
@ -171,9 +177,6 @@ class GraphRag:
tpl_request_queue=None, tpl_request_queue=None,
tpl_response_queue=None, tpl_response_queue=None,
verbose=False, verbose=False,
entity_limit=50,
triple_limit=30,
max_subgraph_size=3000,
module="test", module="test",
): ):
@ -230,10 +233,6 @@ class GraphRag:
subscriber=module + "-emb", subscriber=module + "-emb",
) )
self.entity_limit=entity_limit
self.query_limit=triple_limit
self.max_subgraph_size=max_subgraph_size
self.label_cache = {} self.label_cache = {}
self.prompt = PromptClient( self.prompt = PromptClient(
@ -247,13 +246,18 @@ class GraphRag:
if self.verbose: if self.verbose:
print("Initialised", flush=True) print("Initialised", flush=True)
def query(self, query, user="trustgraph", collection="default"): def query(
self, query, user="trustgraph", collection="default",
entity_limit=50, triple_limit=30, max_subgraph_size=1000,
):
if self.verbose: if self.verbose:
print("Construct prompt...", flush=True) print("Construct prompt...", flush=True)
q = Query( q = Query(
rag=self, user=user, collection=collection, verbose=self.verbose rag=self, user=user, collection=collection, verbose=self.verbose
entity_limit=entity_limit, query_limit=query_limit,
max_subgraph_size=max_subgraph_size,
) )
kg = q.get_labelgraph(query) kg = q.get_labelgraph(query)

View file

@ -50,6 +50,8 @@ class Processor(ConsumerProducer):
document_embeddings_response_queue document_embeddings_response_queue
) )
doc_limit = params.get("doc_limit", 10)
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue, "input_queue": input_queue,
@ -79,6 +81,8 @@ class Processor(ConsumerProducer):
module=module, module=module,
) )
self.doc_limit = doc_limit
async def handle(self, msg): async def handle(self, msg):
try: try:
@ -90,7 +94,12 @@ class Processor(ConsumerProducer):
print(f"Handling input {id}...", flush=True) print(f"Handling input {id}...", flush=True)
response = self.rag.query(v.query) if v.doc_limit:
doc_limit = v.doc_limit
else:
doc_limit = self.doc_limit
response = self.rag.query(v.query, doc_limit=doc_limit)
print("Send response...", flush=True) print("Send response...", flush=True)
r = DocumentRagResponse(response = response, error=None) r = DocumentRagResponse(response = response, error=None)
@ -124,6 +133,13 @@ class Processor(ConsumerProducer):
default_output_queue, default_output_queue,
) )
parser.add_argument(
'-d', '--doc-limit',
type=int,
default=20,
help=f'Default document fetch limit (default: 10)'
)
parser.add_argument( parser.add_argument(
'--prompt-request-queue', '--prompt-request-queue',
default=prompt_request_queue, default=prompt_request_queue,

View file

@ -31,9 +31,7 @@ class Processor(ConsumerProducer):
input_queue = params.get("input_queue", default_input_queue) input_queue = params.get("input_queue", default_input_queue)
output_queue = params.get("output_queue", default_output_queue) output_queue = params.get("output_queue", default_output_queue)
subscriber = params.get("subscriber", default_subscriber) subscriber = params.get("subscriber", default_subscriber)
entity_limit = params.get("entity_limit", 50)
triple_limit = params.get("triple_limit", 30)
max_subgraph_size = params.get("max_subgraph_size", 3000)
pr_request_queue = params.get( pr_request_queue = params.get(
"prompt_request_queue", prompt_request_queue "prompt_request_queue", prompt_request_queue
) )
@ -59,6 +57,10 @@ class Processor(ConsumerProducer):
"triples_response_queue", triples_response_queue "triples_response_queue", triples_response_queue
) )
entity_limit = params.get("entity_limit", 50)
triple_limit = params.get("triple_limit", 30)
max_subgraph_size = params.get("max_subgraph_size", 1000)
super(Processor, self).__init__( super(Processor, self).__init__(
**params | { **params | {
"input_queue": input_queue, "input_queue": input_queue,
@ -92,12 +94,13 @@ class Processor(ConsumerProducer):
tpl_request_queue=triples_request_queue, tpl_request_queue=triples_request_queue,
tpl_response_queue=triples_response_queue, tpl_response_queue=triples_response_queue,
verbose=True, verbose=True,
entity_limit=entity_limit,
triple_limit=triple_limit,
max_subgraph_size=max_subgraph_size,
module=module, module=module,
) )
self.default_entity_limit = entity_limit
self.default_triple_limit = triple_limit
self.default_max_subgraph_size = max_subgraph_size
async def handle(self, msg): async def handle(self, msg):
try: try:
@ -106,15 +109,32 @@ class Processor(ConsumerProducer):
# Sender-produced ID # Sender-produced ID
id = msg.properties()["id"] id = msg.properties()["id"]
print(f"Handling input {id}...", flush=True) print(f"Handling input {id}...", flush=True)
if v.entity_limit:
entity_limit = v.entity_limit
else:
entity_limit = self.entity_limit
if v.triple_limit:
triple_limit = v.triple_limit
else:
triple_limit = self.triple_limit
if v.max_subgraph_size:
max_subgraph_size = v.max_subgraph_size
else:
max_subgraph_size = self.max_subgraph_size
response = self.rag.query( response = self.rag.query(
query=v.query, user=v.user, collection=v.collection query=v.query, user=v.user, collection=v.collection,
entity_limit=entity_limit, triple_limit=triple_limit,
max_subgraph_size=max_subgraph_size
) )
print("Send response...", flush=True) print("Send response...", flush=True)
r = GraphRagResponse(response = response, error=None) r = GraphRagResponse(response=response, error=None)
await self.send(r, properties={"id": id}) await self.send(r, properties={"id": id})
print("Done.", flush=True) print("Done.", flush=True)
@ -149,21 +169,21 @@ class Processor(ConsumerProducer):
'-e', '--entity-limit', '-e', '--entity-limit',
type=int, type=int,
default=50, default=50,
help=f'Entity vector fetch limit (default: 50)' help=f'Default entity vector fetch limit (default: 50)'
) )
parser.add_argument( parser.add_argument(
'-t', '--triple-limit', '-t', '--triple-limit',
type=int, type=int,
default=30, default=30,
help=f'Triple query limit, per query (default: 30)' help=f'Default triple query limit, per query (default: 30)'
) )
parser.add_argument( parser.add_argument(
'-u', '--max-subgraph-size', '-u', '--max-subgraph-size',
type=int, type=int,
default=3000, default=1000,
help=f'Max subgraph size (default: 3000)' help=f'Default max subgraph size (default: 1000)'
) )
parser.add_argument( parser.add_argument(