Message tree state machine (#555)

* add query_incomplete_rankings()

* Add SQL queries for TreeManager task selection

* first working version of TreeManager.next_task()

* remove old generate_task(), add mandatory_labels to text_labels task

* Add ConversationMessage list to Ranking tasks

* add more sophisticated sql queries to find extendible trees

* add TreeManager.query_extendible_parents()

* fix task validation, seed data insertion (reviewed)

* provide user for task selection in text-frontend

* enter 'growing' state

* enter 'aborted_low_grade' state

* enter 'ranking' state

* check tree 'growing' state upon relpy insertion

* exclude user from labeling their own messages (added DEBUG_ALLOW_SELF_LABELING setting)

* add DEBUG_ALLOW_SELF_LABELING to docker-compose.yaml

* fix ranking submission

* add query_tree_ranking_results()

* add ranked_message_ids to RankingReactionPayload

* fix reply_messages instead of prompt_messages

* incorment 'ranking_count' of ranked replies

* added logic to check_condition_for_scoring_state

* changes to msg_tree_state_machine

* pre-commit changes

* enter 'ready_for_scoring' state

* re-add HF embedding call (lost during merge)

* use prepare_conversation() helper for seed-data creation

* Partially add user specified task selection

Co-authored-by: Daniel Hug <danielpatrickhug@gmail.com>
This commit is contained in:
Andreas Köpf
2023-01-11 10:54:03 +01:00
committed by GitHub
co-authored by Daniel Hug
parent 23ff01c603
commit 14fa08e2e7
19 changed files with 1212 additions and 323 deletions
+29 -13
View File
@@ -12,9 +12,12 @@ from fastapi_limiter import FastAPILimiter
from loguru import logger
from oasst_backend.api.deps import get_dummy_api_client
from oasst_backend.api.v1.api import api_router
from oasst_backend.api.v1.utils import prepare_conversation
from oasst_backend.config import settings
from oasst_backend.database import engine
from oasst_backend.models import message_tree_state
from oasst_backend.prompt_repository import PromptRepository, TaskRepository, UserRepository
from oasst_backend.tree_manager import TreeManager, TreeManagerConfiguration
from oasst_shared.exceptions import OasstError, OasstErrorCode
from oasst_shared.schemas import protocol as protocol_schema
from pydantic import BaseModel
@@ -116,6 +119,7 @@ if settings.DEBUG_USE_SEED_DATA:
pr = PromptRepository(
db=db, api_client=api_client, client_user=dummy_user, user_repository=ur, task_repository=tr
)
tm = TreeManager(db, pr, TreeManagerConfiguration())
with open(settings.DEBUG_USE_SEED_DATA_PATH) as f:
dummy_messages_raw = json.load(f)
@@ -138,24 +142,19 @@ if settings.DEBUG_USE_SEED_DATA:
msg.parent_message_id, fail_if_missing=True
)
conversation_messages = pr.fetch_message_conversation(parent_message)
conversation = protocol_schema.Conversation(
messages=[
protocol_schema.ConversationMessage(
text=cmsg.text,
is_assistant=cmsg.role == "assistant",
message_id=cmsg.id,
fronend_message_id=cmsg.frontend_message_id,
)
for cmsg in conversation_messages
]
)
conversation = prepare_conversation(conversation_messages)
task = tr.store_task(
protocol_schema.AssistantReplyTask(conversation=conversation),
message_tree_id=parent_message.message_tree_id,
parent_message_id=parent_message.id,
)
tr.bind_frontend_message_id(task.id, msg.task_message_id)
message = pr.store_text_reply(msg.text, msg.task_message_id, msg.user_message_id)
message = pr.store_text_reply(
msg.text, msg.task_message_id, msg.user_message_id, review_count=5, review_result=True
)
if message.parent_id is None:
tm._insert_default_state(root_message_id=message.id, state=message_tree_state.State.GROWING)
db.commit()
logger.info(
f"Inserted: message_id: {message.id}, payload: {message.payload.payload}, parent_message_id: {message.parent_id}"
@@ -168,6 +167,19 @@ if settings.DEBUG_USE_SEED_DATA:
logger.exception("Seed data insertion failed")
@app.on_event("startup")
def ensure_tree_states():
try:
logger.info("Startup: TreeManager.ensure_tree_states()")
cfg = TreeManagerConfiguration() # TODO: decide where config is stored, e.g. load form json/yaml file
with Session(engine) as db:
tm = TreeManager(db, None, configuration=cfg)
tm.ensure_tree_states()
except Exception:
logger.exception("TreeManager.ensure_tree_states() failed.")
app.include_router(api_router, prefix=settings.API_V1_STR)
@@ -175,7 +187,7 @@ def get_openapi_schema():
return json.dumps(app.openapi())
if __name__ == "__main__":
def main():
# Importing here so we don't import packages unnecessarily if we're
# importing main as a module.
import argparse
@@ -198,3 +210,7 @@ if __name__ == "__main__":
print(get_openapi_schema())
else:
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()