mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-08-02 12:20:35 +08:00
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:
co-authored by
Daniel Hug
parent
23ff01c603
commit
14fa08e2e7
+29
-13
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user